From d3417cc29f208d3b98254cbc5e54a324b5503678 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Fri, 26 Jun 2026 15:56:22 +0530
Subject: [PATCH 01/14] add: Enabled option to add the toleration while model
deployment.
---
api/api/validator.go | 36 ++++
api/api/validator_test.go | 184 +++++++++++++++++
api/client/model_resource_request.go | 49 ++++-
api/client/model_toleration.go | 263 +++++++++++++++++++++++++
api/cluster/resource/templater.go | 10 +
api/cluster/resource/templater_test.go | 213 ++++++++++++++++++++
api/models/resource_request.go | 3 +
swagger.yaml | 25 +++
8 files changed, 776 insertions(+), 7 deletions(-)
create mode 100644 api/api/validator_test.go
create mode 100644 api/client/model_toleration.go
diff --git a/api/api/validator.go b/api/api/validator.go
index b75cbc17a..9d05f08fb 100644
--- a/api/api/validator.go
+++ b/api/api/validator.go
@@ -6,6 +6,7 @@ import (
"fmt"
"golang.org/x/exp/slices"
+ corev1 "k8s.io/api/core/v1"
"github.com/caraml-dev/merlin/config"
"github.com/caraml-dev/merlin/models"
@@ -73,10 +74,45 @@ func resourceRequestValidation(endpoint *models.VersionEndpoint) requestValidato
return fmt.Errorf("max replica must be greater than 0")
}
+ if err := validateTolerations(endpoint.ResourceRequest.Tolerations); err != nil {
+ return fmt.Errorf("invalid toleration in resource request: %w", err)
+ }
+
return nil
})
}
+var validTolerationEffects = []corev1.TaintEffect{
+ corev1.TaintEffectNoSchedule,
+ corev1.TaintEffectPreferNoSchedule,
+ corev1.TaintEffectNoExecute,
+ "", // empty matches all effects
+}
+
+var validTolerationOperators = []corev1.TolerationOperator{
+ corev1.TolerationOpEqual,
+ corev1.TolerationOpExists,
+ "", // empty defaults to Equal
+}
+
+func validateTolerations(tolerations []corev1.Toleration) error {
+ for i, t := range tolerations {
+ if !slices.Contains(validTolerationEffects, t.Effect) {
+ return fmt.Errorf("toleration[%d] has invalid effect %q; must be one of: NoSchedule, PreferNoSchedule, NoExecute", i, t.Effect)
+ }
+ if !slices.Contains(validTolerationOperators, t.Operator) {
+ return fmt.Errorf("toleration[%d] has invalid operator %q; must be one of: Equal, Exists", i, t.Operator)
+ }
+ if t.Operator == corev1.TolerationOpExists && t.Value != "" {
+ return fmt.Errorf("toleration[%d] with operator 'Exists' must not specify a value", i)
+ }
+ if t.TolerationSeconds != nil && t.Effect != corev1.TaintEffectNoExecute {
+ return fmt.Errorf("toleration[%d] tolerationSeconds is only valid for effect 'NoExecute'", i)
+ }
+ }
+ return nil
+}
+
func customModelValidation(model *models.Model, version *models.Version) requestValidator {
return newFuncValidate(func() error {
if model.Type == models.ModelTypeCustom {
diff --git a/api/api/validator_test.go b/api/api/validator_test.go
new file mode 100644
index 000000000..23d1000af
--- /dev/null
+++ b/api/api/validator_test.go
@@ -0,0 +1,184 @@
+package api
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/assert"
+ corev1 "k8s.io/api/core/v1"
+
+ "github.com/caraml-dev/merlin/models"
+)
+
+func int64Ptr(v int64) *int64 { return &v }
+
+func TestValidateTolerations(t *testing.T) {
+ tests := []struct {
+ name string
+ tolerations []corev1.Toleration
+ wantErr bool
+ errContains string
+ }{
+ {
+ name: "nil tolerations — always valid",
+ tolerations: nil,
+ wantErr: false,
+ },
+ {
+ name: "empty tolerations — always valid",
+ tolerations: []corev1.Toleration{},
+ wantErr: false,
+ },
+ {
+ name: "valid Equal toleration with NoSchedule",
+ tolerations: []corev1.Toleration{
+ {Key: "dedicated", Operator: corev1.TolerationOpEqual, Value: "ml-team", Effect: corev1.TaintEffectNoSchedule},
+ },
+ wantErr: false,
+ },
+ {
+ name: "valid Exists toleration with empty value and NoExecute",
+ tolerations: []corev1.Toleration{
+ {Key: "spot", Operator: corev1.TolerationOpExists, Effect: corev1.TaintEffectNoExecute},
+ },
+ wantErr: false,
+ },
+ {
+ name: "valid empty operator defaults to Equal",
+ tolerations: []corev1.Toleration{
+ {Key: "key", Operator: "", Value: "val", Effect: corev1.TaintEffectPreferNoSchedule},
+ },
+ wantErr: false,
+ },
+ {
+ name: "valid empty effect matches all effects",
+ tolerations: []corev1.Toleration{
+ {Key: "key", Operator: corev1.TolerationOpEqual, Value: "val", Effect: ""},
+ },
+ wantErr: false,
+ },
+ {
+ name: "valid TolerationSeconds on NoExecute",
+ tolerations: []corev1.Toleration{
+ {Key: "key", Operator: corev1.TolerationOpExists, Effect: corev1.TaintEffectNoExecute, TolerationSeconds: int64Ptr(300)},
+ },
+ wantErr: false,
+ },
+ {
+ name: "multiple valid tolerations",
+ tolerations: []corev1.Toleration{
+ {Key: "key1", Operator: corev1.TolerationOpEqual, Value: "v1", Effect: corev1.TaintEffectNoSchedule},
+ {Key: "key2", Operator: corev1.TolerationOpExists, Effect: corev1.TaintEffectNoExecute, TolerationSeconds: int64Ptr(60)},
+ },
+ wantErr: false,
+ },
+ // --- invalid cases ---
+ {
+ name: "invalid effect",
+ tolerations: []corev1.Toleration{
+ {Key: "key", Operator: corev1.TolerationOpEqual, Value: "val", Effect: "InvalidEffect"},
+ },
+ wantErr: true,
+ errContains: "invalid effect",
+ },
+ {
+ name: "invalid operator",
+ tolerations: []corev1.Toleration{
+ {Key: "key", Operator: "NotAValidOperator", Value: "val", Effect: corev1.TaintEffectNoSchedule},
+ },
+ wantErr: true,
+ errContains: "invalid operator",
+ },
+ {
+ name: "Exists operator with a non-empty value",
+ tolerations: []corev1.Toleration{
+ {Key: "key", Operator: corev1.TolerationOpExists, Value: "should-be-empty", Effect: corev1.TaintEffectNoSchedule},
+ },
+ wantErr: true,
+ errContains: "must not specify a value",
+ },
+ {
+ name: "TolerationSeconds on NoSchedule (only valid for NoExecute)",
+ tolerations: []corev1.Toleration{
+ {Key: "key", Operator: corev1.TolerationOpEqual, Value: "v", Effect: corev1.TaintEffectNoSchedule, TolerationSeconds: int64Ptr(120)},
+ },
+ wantErr: true,
+ errContains: "tolerationSeconds is only valid for effect 'NoExecute'",
+ },
+ {
+ name: "TolerationSeconds on PreferNoSchedule",
+ tolerations: []corev1.Toleration{
+ {Key: "key", Operator: corev1.TolerationOpEqual, Value: "v", Effect: corev1.TaintEffectPreferNoSchedule, TolerationSeconds: int64Ptr(60)},
+ },
+ wantErr: true,
+ errContains: "tolerationSeconds is only valid for effect 'NoExecute'",
+ },
+ {
+ name: "second toleration in list is invalid",
+ tolerations: []corev1.Toleration{
+ {Key: "ok", Operator: corev1.TolerationOpEqual, Value: "v", Effect: corev1.TaintEffectNoSchedule},
+ {Key: "bad", Operator: corev1.TolerationOpExists, Value: "oops", Effect: corev1.TaintEffectNoSchedule},
+ },
+ wantErr: true,
+ errContains: "must not specify a value",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ err := validateTolerations(tt.tolerations)
+ if tt.wantErr {
+ assert.Error(t, err)
+ if tt.errContains != "" {
+ assert.Contains(t, err.Error(), tt.errContains)
+ }
+ } else {
+ assert.NoError(t, err)
+ }
+ })
+ }
+}
+
+func TestResourceRequestValidation_Tolerations(t *testing.T) {
+ tests := []struct {
+ name string
+ req *models.ResourceRequest
+ wantErr bool
+ }{
+ {
+ name: "nil resource request passes",
+ req: nil,
+ },
+ {
+ name: "valid resource request with tolerations passes",
+ req: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ Tolerations: []corev1.Toleration{
+ {Key: "dedicated", Operator: corev1.TolerationOpEqual, Value: "ml", Effect: corev1.TaintEffectNoSchedule},
+ },
+ },
+ },
+ {
+ name: "resource request with invalid toleration operator fails",
+ req: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ Tolerations: []corev1.Toleration{
+ {Key: "k", Operator: "BadOp", Value: "v"},
+ },
+ },
+ wantErr: true,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ endpoint := &models.VersionEndpoint{ResourceRequest: tt.req}
+ validator := resourceRequestValidation(endpoint)
+ err := validator.validate()
+ if tt.wantErr {
+ assert.Error(t, err)
+ } else {
+ assert.NoError(t, err)
+ }
+ })
+ }
+}
diff --git a/api/client/model_resource_request.go b/api/client/model_resource_request.go
index ccd65559a..bd719c229 100644
--- a/api/client/model_resource_request.go
+++ b/api/client/model_resource_request.go
@@ -19,13 +19,14 @@ var _ MappedNullable = &ResourceRequest{}
// ResourceRequest struct for ResourceRequest
type ResourceRequest struct {
- MinReplica *int32 `json:"min_replica,omitempty"`
- MaxReplica *int32 `json:"max_replica,omitempty"`
- CpuRequest *string `json:"cpu_request,omitempty"`
- CpuLimit *string `json:"cpu_limit,omitempty"`
- MemoryRequest *string `json:"memory_request,omitempty"`
- GpuName *string `json:"gpu_name,omitempty"`
- GpuRequest *string `json:"gpu_request,omitempty"`
+ MinReplica *int32 `json:"min_replica,omitempty"`
+ MaxReplica *int32 `json:"max_replica,omitempty"`
+ CpuRequest *string `json:"cpu_request,omitempty"`
+ CpuLimit *string `json:"cpu_limit,omitempty"`
+ MemoryRequest *string `json:"memory_request,omitempty"`
+ GpuName *string `json:"gpu_name,omitempty"`
+ GpuRequest *string `json:"gpu_request,omitempty"`
+ Tolerations []Toleration `json:"tolerations,omitempty"`
}
// NewResourceRequest instantiates a new ResourceRequest object
@@ -269,6 +270,37 @@ func (o *ResourceRequest) SetGpuRequest(v string) {
o.GpuRequest = &v
}
+// GetTolerations returns the Tolerations field value if set, zero value otherwise.
+func (o *ResourceRequest) GetTolerations() []Toleration {
+ if o == nil || IsNil(o.Tolerations) {
+ var ret []Toleration
+ return ret
+ }
+ return o.Tolerations
+}
+
+// GetTolerationsOk returns a tuple with the Tolerations field value if set, nil otherwise
+// and a boolean to check if the value has been set.
+func (o *ResourceRequest) GetTolerationsOk() ([]Toleration, bool) {
+ if o == nil || IsNil(o.Tolerations) {
+ return nil, false
+ }
+ return o.Tolerations, true
+}
+
+// HasTolerations returns a boolean if a field has been set.
+func (o *ResourceRequest) HasTolerations() bool {
+ if o != nil && !IsNil(o.Tolerations) {
+ return true
+ }
+ return false
+}
+
+// SetTolerations gets a reference to the given []Toleration and assigns it to the Tolerations field.
+func (o *ResourceRequest) SetTolerations(v []Toleration) {
+ o.Tolerations = v
+}
+
func (o ResourceRequest) MarshalJSON() ([]byte, error) {
toSerialize, err := o.ToMap()
if err != nil {
@@ -300,6 +332,9 @@ func (o ResourceRequest) ToMap() (map[string]interface{}, error) {
if !IsNil(o.GpuRequest) {
toSerialize["gpu_request"] = o.GpuRequest
}
+ if !IsNil(o.Tolerations) {
+ toSerialize["tolerations"] = o.Tolerations
+ }
return toSerialize, nil
}
diff --git a/api/client/model_toleration.go b/api/client/model_toleration.go
new file mode 100644
index 000000000..ad31e9eaa
--- /dev/null
+++ b/api/client/model_toleration.go
@@ -0,0 +1,263 @@
+/*
+Merlin
+
+API Guide for accessing Merlin's model management, deployment, and serving functionalities
+
+API version: 0.14.0
+*/
+
+// Code generated by OpenAPI Generator (https://openapi-generator.tech); DO NOT EDIT.
+
+package client
+
+import (
+ "encoding/json"
+)
+
+// checks if the Toleration type satisfies the MappedNullable interface at compile time
+var _ MappedNullable = &Toleration{}
+
+// Toleration Kubernetes toleration for scheduling pods onto tainted nodes
+type Toleration struct {
+ // The taint key the toleration applies to
+ Key *string `json:"key,omitempty"`
+ // 'Exists' or 'Equal' (default: Equal)
+ Operator *string `json:"operator,omitempty"`
+ // The taint value the toleration matches (only used when operator is Equal)
+ Value *string `json:"value,omitempty"`
+ // 'NoSchedule', 'PreferNoSchedule', or 'NoExecute'. Empty matches all effects
+ Effect *string `json:"effect,omitempty"`
+ // Period (seconds) the pod tolerates the taint (only for NoExecute effect)
+ TolerationSeconds *int64 `json:"toleration_seconds,omitempty"`
+}
+
+// NewToleration instantiates a new Toleration object
+func NewToleration() *Toleration {
+ this := Toleration{}
+ return &this
+}
+
+// NewTolerationWithDefaults instantiates a new Toleration object with defaults
+func NewTolerationWithDefaults() *Toleration {
+ this := Toleration{}
+ return &this
+}
+
+// GetKey returns the Key field value if set, zero value otherwise.
+func (o *Toleration) GetKey() string {
+ if o == nil || IsNil(o.Key) {
+ var ret string
+ return ret
+ }
+ return *o.Key
+}
+
+// GetKeyOk returns a tuple with the Key field value if set, nil otherwise
+// and a boolean to check if the value has been set.
+func (o *Toleration) GetKeyOk() (*string, bool) {
+ if o == nil || IsNil(o.Key) {
+ return nil, false
+ }
+ return o.Key, true
+}
+
+// HasKey returns a boolean if a field has been set.
+func (o *Toleration) HasKey() bool {
+ if o != nil && !IsNil(o.Key) {
+ return true
+ }
+ return false
+}
+
+// SetKey gets a reference to the given string and assigns it to the Key field.
+func (o *Toleration) SetKey(v string) {
+ o.Key = &v
+}
+
+// GetOperator returns the Operator field value if set, zero value otherwise.
+func (o *Toleration) GetOperator() string {
+ if o == nil || IsNil(o.Operator) {
+ var ret string
+ return ret
+ }
+ return *o.Operator
+}
+
+// GetOperatorOk returns a tuple with the Operator field value if set, nil otherwise
+// and a boolean to check if the value has been set.
+func (o *Toleration) GetOperatorOk() (*string, bool) {
+ if o == nil || IsNil(o.Operator) {
+ return nil, false
+ }
+ return o.Operator, true
+}
+
+// HasOperator returns a boolean if a field has been set.
+func (o *Toleration) HasOperator() bool {
+ if o != nil && !IsNil(o.Operator) {
+ return true
+ }
+ return false
+}
+
+// SetOperator gets a reference to the given string and assigns it to the Operator field.
+func (o *Toleration) SetOperator(v string) {
+ o.Operator = &v
+}
+
+// GetValue returns the Value field value if set, zero value otherwise.
+func (o *Toleration) GetValue() string {
+ if o == nil || IsNil(o.Value) {
+ var ret string
+ return ret
+ }
+ return *o.Value
+}
+
+// GetValueOk returns a tuple with the Value field value if set, nil otherwise
+// and a boolean to check if the value has been set.
+func (o *Toleration) GetValueOk() (*string, bool) {
+ if o == nil || IsNil(o.Value) {
+ return nil, false
+ }
+ return o.Value, true
+}
+
+// HasValue returns a boolean if a field has been set.
+func (o *Toleration) HasValue() bool {
+ if o != nil && !IsNil(o.Value) {
+ return true
+ }
+ return false
+}
+
+// SetValue gets a reference to the given string and assigns it to the Value field.
+func (o *Toleration) SetValue(v string) {
+ o.Value = &v
+}
+
+// GetEffect returns the Effect field value if set, zero value otherwise.
+func (o *Toleration) GetEffect() string {
+ if o == nil || IsNil(o.Effect) {
+ var ret string
+ return ret
+ }
+ return *o.Effect
+}
+
+// GetEffectOk returns a tuple with the Effect field value if set, nil otherwise
+// and a boolean to check if the value has been set.
+func (o *Toleration) GetEffectOk() (*string, bool) {
+ if o == nil || IsNil(o.Effect) {
+ return nil, false
+ }
+ return o.Effect, true
+}
+
+// HasEffect returns a boolean if a field has been set.
+func (o *Toleration) HasEffect() bool {
+ if o != nil && !IsNil(o.Effect) {
+ return true
+ }
+ return false
+}
+
+// SetEffect gets a reference to the given string and assigns it to the Effect field.
+func (o *Toleration) SetEffect(v string) {
+ o.Effect = &v
+}
+
+// GetTolerationSeconds returns the TolerationSeconds field value if set, zero value otherwise.
+func (o *Toleration) GetTolerationSeconds() int64 {
+ if o == nil || IsNil(o.TolerationSeconds) {
+ var ret int64
+ return ret
+ }
+ return *o.TolerationSeconds
+}
+
+// GetTolerationSecondsOk returns a tuple with the TolerationSeconds field value if set, nil otherwise
+// and a boolean to check if the value has been set.
+func (o *Toleration) GetTolerationSecondsOk() (*int64, bool) {
+ if o == nil || IsNil(o.TolerationSeconds) {
+ return nil, false
+ }
+ return o.TolerationSeconds, true
+}
+
+// HasTolerationSeconds returns a boolean if a field has been set.
+func (o *Toleration) HasTolerationSeconds() bool {
+ if o != nil && !IsNil(o.TolerationSeconds) {
+ return true
+ }
+ return false
+}
+
+// SetTolerationSeconds gets a reference to the given int64 and assigns it to the TolerationSeconds field.
+func (o *Toleration) SetTolerationSeconds(v int64) {
+ o.TolerationSeconds = &v
+}
+
+func (o Toleration) MarshalJSON() ([]byte, error) {
+ toSerialize, err := o.ToMap()
+ if err != nil {
+ return []byte{}, err
+ }
+ return json.Marshal(toSerialize)
+}
+
+func (o Toleration) ToMap() (map[string]interface{}, error) {
+ toSerialize := map[string]interface{}{}
+ if !IsNil(o.Key) {
+ toSerialize["key"] = o.Key
+ }
+ if !IsNil(o.Operator) {
+ toSerialize["operator"] = o.Operator
+ }
+ if !IsNil(o.Value) {
+ toSerialize["value"] = o.Value
+ }
+ if !IsNil(o.Effect) {
+ toSerialize["effect"] = o.Effect
+ }
+ if !IsNil(o.TolerationSeconds) {
+ toSerialize["toleration_seconds"] = o.TolerationSeconds
+ }
+ return toSerialize, nil
+}
+
+type NullableToleration struct {
+ value *Toleration
+ isSet bool
+}
+
+func (v NullableToleration) Get() *Toleration {
+ return v.value
+}
+
+func (v *NullableToleration) Set(val *Toleration) {
+ v.value = val
+ v.isSet = true
+}
+
+func (v NullableToleration) IsSet() bool {
+ return v.isSet
+}
+
+func (v *NullableToleration) Unset() {
+ v.value = nil
+ v.isSet = false
+}
+
+func NewNullableToleration(val *Toleration) *NullableToleration {
+ return &NullableToleration{value: val, isSet: true}
+}
+
+func (v NullableToleration) MarshalJSON() ([]byte, error) {
+ return json.Marshal(v.value)
+}
+
+func (v *NullableToleration) UnmarshalJSON(src []byte) error {
+ v.isSet = true
+ return json.Unmarshal(src, &v.value)
+}
diff --git a/api/cluster/resource/templater.go b/api/cluster/resource/templater.go
index 3018e1bc6..06864f559 100644
--- a/api/cluster/resource/templater.go
+++ b/api/cluster/resource/templater.go
@@ -237,6 +237,11 @@ func (t *InferenceServiceTemplater) createPredictorSpec(modelService *models.Ser
}
}
+ // Append user-defined tolerations from ResourceRequest (merged on top of any GPU-derived tolerations)
+ if len(modelService.ResourceRequest.Tolerations) > 0 {
+ tolerations = append(tolerations, modelService.ResourceRequest.Tolerations...)
+ }
+
// Get user-configured probe settings
var userLivenessConfig *models.ProbeConfig
var userReadinessConfig *models.ProbeConfig
@@ -468,6 +473,11 @@ func (t *InferenceServiceTemplater) createTransformerSpec(
},
}
+ // Apply user-defined tolerations for transformer pods
+ if len(transformer.ResourceRequest.Tolerations) > 0 {
+ transformerSpec.PodSpec.Tolerations = transformer.ResourceRequest.Tolerations
+ }
+
return transformerSpec, nil
}
diff --git a/api/cluster/resource/templater_test.go b/api/cluster/resource/templater_test.go
index ff0e9e620..91ac5afb5 100644
--- a/api/cluster/resource/templater_test.go
+++ b/api/cluster/resource/templater_test.go
@@ -4794,3 +4794,216 @@ func sortInferenceServiceSpecEnvVars(isvc kservev1beta1.InferenceServiceSpec) {
}
}
}
+
+func TestCreateInferenceServiceSpecWithTolerations(t *testing.T) {
+ err := labeller.InitKubernetesLabeller("gojek.com/", "caraml.dev/", testEnvironmentName)
+ assert.NoError(t, err)
+ defer func() { _ = labeller.InitKubernetesLabeller("", "", "") }()
+
+ project := mlp.Project{Name: "project"}
+
+ userTolerations := []corev1.Toleration{
+ {
+ Key: "dedicated",
+ Operator: corev1.TolerationOpEqual,
+ Value: "ml-team",
+ Effect: corev1.TaintEffectNoSchedule,
+ },
+ }
+
+ gpuTolerations := []corev1.Toleration{
+ {
+ Key: "nvidia.com/gpu",
+ Operator: corev1.TolerationOpExists,
+ Effect: corev1.TaintEffectNoSchedule,
+ },
+ }
+
+ gpuConfig := config.GPUConfig{
+ Name: "NVIDIA T4",
+ Values: []string{"1"},
+ ResourceType: "nvidia.com/gpu",
+ NodeSelector: map[string]string{"cloud.google.com/gke-accelerator": "nvidia-tesla-t4"},
+ Tolerations: gpuTolerations,
+ }
+
+ baseModelSvc := &models.Service{
+ Name: "model-1",
+ ModelName: "model",
+ Namespace: "project",
+ ModelVersion: "1",
+ ArtifactURI: "gs://my-artifacet",
+ Metadata: models.Metadata{
+ App: "model",
+ Component: models.ComponentModelVersion,
+ Stream: "dsp",
+ Team: "dsp",
+ },
+ Protocol: protocol.HttpJson,
+ }
+
+ baseDeployConfig := &config.DeploymentConfig{
+ DefaultModelResourceRequests: defaultModelResourceRequests,
+ DefaultTransformerResourceRequests: defaultTransformerResourceRequests,
+ QueueResourcePercentage: "2",
+ StandardTransformer: standardTransformerConfig,
+ UserContainerCPUDefaultLimit: userContainerCPUDefaultLimit,
+ UserContainerCPULimitRequestFactor: userContainerCPULimitRequestFactor,
+ UserContainerMemoryLimitRequestFactor: userContainerMemoryLimitRequestFactor,
+ DefaultEnvVarsWithoutCPULimits: []corev1.EnvVar{defaultEnvVarWithoutCPULimits},
+ }
+
+ tests := []struct {
+ name string
+ modelSvc *models.Service
+ deployConfig *config.DeploymentConfig
+ wantErr bool
+ checkFn func(t *testing.T, infSvc *kservev1beta1.InferenceService)
+ }{
+ {
+ name: "predictor with user-defined tolerations only",
+ modelSvc: &models.Service{
+ Name: baseModelSvc.Name,
+ ModelName: baseModelSvc.ModelName,
+ ModelVersion: baseModelSvc.ModelVersion,
+ Namespace: project.Name,
+ ArtifactURI: baseModelSvc.ArtifactURI,
+ Type: models.ModelTypeTensorflow,
+ Options: &models.ModelOption{},
+ Metadata: baseModelSvc.Metadata,
+ Protocol: protocol.HttpJson,
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1,
+ MaxReplica: 2,
+ CPURequest: resource.MustParse("500m"),
+ MemoryRequest: resource.MustParse("500Mi"),
+ Tolerations: userTolerations,
+ },
+ },
+ deployConfig: baseDeployConfig,
+ checkFn: func(t *testing.T, infSvc *kservev1beta1.InferenceService) {
+ assert.Equal(t, userTolerations, infSvc.Spec.Predictor.Tolerations,
+ "predictor should carry user-defined tolerations")
+ assert.Nil(t, infSvc.Spec.Transformer)
+ },
+ },
+ {
+ name: "predictor with GPU tolerations merged with user-defined tolerations",
+ modelSvc: &models.Service{
+ Name: baseModelSvc.Name,
+ ModelName: baseModelSvc.ModelName,
+ ModelVersion: baseModelSvc.ModelVersion,
+ Namespace: project.Name,
+ ArtifactURI: baseModelSvc.ArtifactURI,
+ Type: models.ModelTypeTensorflow,
+ Options: &models.ModelOption{},
+ Metadata: baseModelSvc.Metadata,
+ Protocol: protocol.HttpJson,
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1,
+ MaxReplica: 2,
+ CPURequest: resource.MustParse("500m"),
+ MemoryRequest: resource.MustParse("500Mi"),
+ GPUName: "NVIDIA T4",
+ GPURequest: resource.MustParse("1"),
+ Tolerations: userTolerations,
+ },
+ },
+ deployConfig: func() *config.DeploymentConfig {
+ cfg := *baseDeployConfig
+ cfg.GPUs = []config.GPUConfig{gpuConfig}
+ return &cfg
+ }(),
+ checkFn: func(t *testing.T, infSvc *kservev1beta1.InferenceService) {
+ got := infSvc.Spec.Predictor.Tolerations
+ assert.Equal(t, append(gpuTolerations, userTolerations...), got,
+ "predictor should have GPU tolerations followed by user tolerations")
+ },
+ },
+ {
+ name: "predictor with no tolerations has empty toleration list",
+ modelSvc: &models.Service{
+ Name: baseModelSvc.Name,
+ ModelName: baseModelSvc.ModelName,
+ ModelVersion: baseModelSvc.ModelVersion,
+ Namespace: project.Name,
+ ArtifactURI: baseModelSvc.ArtifactURI,
+ Type: models.ModelTypeTensorflow,
+ Options: &models.ModelOption{},
+ Metadata: baseModelSvc.Metadata,
+ Protocol: protocol.HttpJson,
+ },
+ deployConfig: baseDeployConfig,
+ checkFn: func(t *testing.T, infSvc *kservev1beta1.InferenceService) {
+ assert.Empty(t, infSvc.Spec.Predictor.Tolerations)
+ },
+ },
+ {
+ name: "transformer with user-defined tolerations",
+ modelSvc: &models.Service{
+ Name: baseModelSvc.Name,
+ ModelName: baseModelSvc.ModelName,
+ ModelVersion: baseModelSvc.ModelVersion,
+ Namespace: project.Name,
+ ArtifactURI: baseModelSvc.ArtifactURI,
+ Type: models.ModelTypeTensorflow,
+ Options: &models.ModelOption{},
+ Metadata: baseModelSvc.Metadata,
+ Protocol: protocol.HttpJson,
+ Transformer: &models.Transformer{
+ Enabled: true,
+ Image: "ghcr.io/gojek/merlin-transformer-test",
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1,
+ MaxReplica: 2,
+ CPURequest: resource.MustParse("100m"),
+ MemoryRequest: resource.MustParse("500Mi"),
+ Tolerations: userTolerations,
+ },
+ },
+ },
+ deployConfig: baseDeployConfig,
+ checkFn: func(t *testing.T, infSvc *kservev1beta1.InferenceService) {
+ assert.NotNil(t, infSvc.Spec.Transformer)
+ assert.Equal(t, userTolerations, infSvc.Spec.Transformer.PodSpec.Tolerations,
+ "transformer pods should carry user-defined tolerations")
+ },
+ },
+ {
+ name: "transformer without tolerations has no tolerations set",
+ modelSvc: &models.Service{
+ Name: baseModelSvc.Name,
+ ModelName: baseModelSvc.ModelName,
+ ModelVersion: baseModelSvc.ModelVersion,
+ Namespace: project.Name,
+ ArtifactURI: baseModelSvc.ArtifactURI,
+ Type: models.ModelTypeTensorflow,
+ Options: &models.ModelOption{},
+ Metadata: baseModelSvc.Metadata,
+ Protocol: protocol.HttpJson,
+ Transformer: &models.Transformer{
+ Enabled: true,
+ Image: "ghcr.io/gojek/merlin-transformer-test",
+ },
+ },
+ deployConfig: baseDeployConfig,
+ checkFn: func(t *testing.T, infSvc *kservev1beta1.InferenceService) {
+ assert.NotNil(t, infSvc.Spec.Transformer)
+ assert.Empty(t, infSvc.Spec.Transformer.PodSpec.Tolerations)
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ tpl := NewInferenceServiceTemplater(*tt.deployConfig)
+ infSvc, err := tpl.CreateInferenceServiceSpec(tt.modelSvc, defaultDeploymentScale)
+ if tt.wantErr {
+ assert.Error(t, err)
+ return
+ }
+ assert.NoError(t, err)
+ tt.checkFn(t, infSvc)
+ })
+ }
+}
diff --git a/api/models/resource_request.go b/api/models/resource_request.go
index d335c9632..839641908 100644
--- a/api/models/resource_request.go
+++ b/api/models/resource_request.go
@@ -19,6 +19,7 @@ import (
"encoding/json"
"errors"
+ corev1 "k8s.io/api/core/v1"
"k8s.io/apimachinery/pkg/api/resource"
)
@@ -41,6 +42,8 @@ type ResourceRequest struct {
LivenessProbe *ProbeConfig `json:"liveness_probe,omitempty"`
// Readiness probe configuration
ReadinessProbe *ProbeConfig `json:"readiness_probe,omitempty"`
+ // Tolerations allow the model pods to be scheduled onto nodes with matching taints
+ Tolerations []corev1.Toleration `json:"tolerations,omitempty"`
}
// ProbeConfig represents the configuration for Kubernetes liveness/readiness probes
diff --git a/swagger.yaml b/swagger.yaml
index 65bece648..182a99e19 100644
--- a/swagger.yaml
+++ b/swagger.yaml
@@ -2121,6 +2121,31 @@ components:
"$ref": "#/components/schemas/ProbeConfig"
readiness_probe:
"$ref": "#/components/schemas/ProbeConfig"
+ tolerations:
+ type: array
+ description: Tolerations allow model pods to be scheduled onto nodes with matching taints
+ items:
+ "$ref": "#/components/schemas/Toleration"
+ Toleration:
+ type: object
+ description: Kubernetes toleration for scheduling pods onto tainted nodes
+ properties:
+ key:
+ type: string
+ description: The taint key the toleration applies to
+ operator:
+ type: string
+ description: "'Exists' or 'Equal' (default: Equal)"
+ value:
+ type: string
+ description: The taint value the toleration matches (only used when operator is Equal)
+ effect:
+ type: string
+ description: "'NoSchedule', 'PreferNoSchedule', or 'NoExecute'. Empty matches all effects"
+ toleration_seconds:
+ type: integer
+ format: int64
+ description: Period (seconds) the pod tolerates the taint (only for NoExecute effect)
ProbeConfig:
type: object
properties:
From 3682189aef9e25f62139bed00652df60b9e58969 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Fri, 26 Jun 2026 16:28:04 +0530
Subject: [PATCH 02/14] add: made UI changes for toleration.
---
ui/src/components/ResourcesConfigTable.js | 16 ++
.../forms/DeployModelVersionForm.js | 36 ++++
.../forms/components/TolerationFormGroup.js | 177 ++++++++++++++++++
.../components/forms/steps/ModelStep.js | 7 +
.../components/forms/steps/TransformerStep.js | 7 +
5 files changed, 243 insertions(+)
create mode 100644 ui/src/pages/version/components/forms/components/TolerationFormGroup.js
diff --git a/ui/src/components/ResourcesConfigTable.js b/ui/src/components/ResourcesConfigTable.js
index dcf9a95cd..49661a9e0 100644
--- a/ui/src/components/ResourcesConfigTable.js
+++ b/ui/src/components/ResourcesConfigTable.js
@@ -29,6 +29,7 @@ export const ResourcesConfigTable = ({
gpu_request,
liveness_probe,
readiness_probe,
+ tolerations,
},
}) => {
const items = [
@@ -104,6 +105,21 @@ export const ResourcesConfigTable = ({
}
}
+ // Add tolerations if configured
+ if (tolerations && tolerations.length > 0) {
+ tolerations.forEach((t, idx) => {
+ const parts = [];
+ if (t.key) parts.push(`key: ${t.key}`);
+ if (t.operator) parts.push(`op: ${t.operator}`);
+ if (t.value) parts.push(`value: ${t.value}`);
+ if (t.effect) parts.push(`effect: ${t.effect}`);
+ items.push({
+ title: idx === 0 ? "Tolerations" : "",
+ description: parts.join(", ") || "—",
+ });
+ });
+ }
+
return (
t && (t.key || t.operator || t.value || t.effect)
+ );
+ if (filtered.length === 0) {
+ delete versionEndpoint.resource_request.tolerations;
+ } else {
+ // Strip fields that are empty strings so the backend doesn't receive them
+ versionEndpoint.resource_request.tolerations = filtered.map((t) => {
+ const entry = {};
+ if (t.key) entry.key = t.key;
+ if (t.operator) entry.operator = t.operator;
+ if (t.value && t.operator !== "Exists") entry.value = t.value;
+ if (t.effect) entry.effect = t.effect;
+ return entry;
+ });
+ }
+ }
+ if (versionEndpoint?.transformer?.resource_request?.tolerations) {
+ const filtered = versionEndpoint.transformer.resource_request.tolerations.filter(
+ (t) => t && (t.key || t.operator || t.value || t.effect)
+ );
+ if (filtered.length === 0) {
+ delete versionEndpoint.transformer.resource_request.tolerations;
+ } else {
+ versionEndpoint.transformer.resource_request.tolerations = filtered.map((t) => {
+ const entry = {};
+ if (t.key) entry.key = t.key;
+ if (t.operator) entry.operator = t.operator;
+ if (t.value && t.operator !== "Exists") entry.value = t.value;
+ if (t.effect) entry.effect = t.effect;
+ return entry;
+ });
+ }
+ }
submitForm({
body: JSON.stringify({
...versionEndpoint,
diff --git a/ui/src/pages/version/components/forms/components/TolerationFormGroup.js b/ui/src/pages/version/components/forms/components/TolerationFormGroup.js
new file mode 100644
index 000000000..ae472b006
--- /dev/null
+++ b/ui/src/pages/version/components/forms/components/TolerationFormGroup.js
@@ -0,0 +1,177 @@
+import React, { Fragment } from "react";
+import {
+ EuiButtonIcon,
+ EuiDescribedFormGroup,
+ EuiFieldText,
+ EuiFlexGroup,
+ EuiFlexItem,
+ EuiSelect,
+ EuiSpacer,
+ EuiText,
+} from "@elastic/eui";
+import { InMemoryTableForm, useOnChangeHandler } from "@caraml-dev/ui-lib";
+
+const OPERATOR_OPTIONS = [
+ { value: "", text: "—" },
+ { value: "Equal", text: "Equal" },
+ { value: "Exists", text: "Exists" },
+];
+
+const EFFECT_OPTIONS = [
+ { value: "", text: "— (Any)" },
+ { value: "NoSchedule", text: "NoSchedule" },
+ { value: "PreferNoSchedule", text: "PreferNoSchedule" },
+ { value: "NoExecute", text: "NoExecute" },
+];
+
+/**
+ * TolerationFormGroup - renders an inline table for adding/removing Kubernetes tolerations.
+ *
+ * Props
+ * tolerations – array of toleration objects (may be undefined / [])
+ * onChangeHandler – function to call when the list changes
+ * errors – validation errors (optional)
+ */
+export const TolerationFormGroup = ({
+ tolerations = [],
+ onChangeHandler,
+ errors = {},
+}) => {
+ const { onChange } = useOnChangeHandler(onChangeHandler);
+
+ const items = [
+ ...tolerations.map((t, idx) => ({ idx, ...t })),
+ { idx: tolerations.length }, // empty "add" row
+ ];
+
+ const onDeleteToleration = (idx) => () => {
+ const updated = [...tolerations];
+ updated.splice(idx, 1);
+ onChangeHandler(updated);
+ };
+
+ const getRowProps = (item) => {
+ const { idx } = item;
+ const isInvalid = !!errors[idx];
+ return {
+ className: isInvalid ? "euiTableRow--isInvalid" : "",
+ "data-test-subj": `toleration-row-${idx}`,
+ };
+ };
+
+ const columns = [
+ {
+ name: "Key",
+ field: "key",
+ width: "22%",
+ render: (key, item) => (
+ onChange(`${item.idx}.key`)(e.target.value)}
+ />
+ ),
+ },
+ {
+ name: "Operator",
+ field: "operator",
+ width: "16%",
+ render: (operator, item) => (
+ {
+ const op = e.target.value;
+ // When switching to Exists, clear the value field
+ if (op === "Exists") {
+ onChange(`${item.idx}.value`)("");
+ }
+ onChange(`${item.idx}.operator`)(op);
+ }}
+ />
+ ),
+ },
+ {
+ name: "Value",
+ field: "value",
+ width: "22%",
+ render: (value, item) => {
+ const isExists = item.operator === "Exists";
+ return (
+ onChange(`${item.idx}.value`)(e.target.value)}
+ />
+ );
+ },
+ },
+ {
+ name: "Effect",
+ field: "effect",
+ width: "26%",
+ render: (effect, item) => (
+ onChange(`${item.idx}.effect`)(e.target.value)}
+ />
+ ),
+ },
+ {
+ width: "14%",
+ actions: [
+ {
+ render: (item) =>
+ item.idx < items.length - 1 ? (
+
+ ) : (
+
+ ),
+ },
+ ],
+ },
+ ];
+
+ return (
+ Node Tolerations
}
+ description={
+
+
+ Allow pods to be scheduled on nodes with matching taints. Leave
+ the last blank row empty to add a new entry.
+
+
+ }
+ fullWidth
+ >
+
+
+
+ `Row ${parseInt(key) + 1}`}
+ />
+
+
+
+ );
+};
+
diff --git a/ui/src/pages/version/components/forms/steps/ModelStep.js b/ui/src/pages/version/components/forms/steps/ModelStep.js
index 4414877c0..7bf977cf8 100644
--- a/ui/src/pages/version/components/forms/steps/ModelStep.js
+++ b/ui/src/pages/version/components/forms/steps/ModelStep.js
@@ -15,6 +15,7 @@ import { ResourcesPanel } from "../components/ResourcesPanel";
import { ImageBuilderSection } from "../components/ImageBuilderSection";
import { CPULimitsFormGroup } from "../components/CPULimitsFormGroup";
import { ProbesFormGroup } from "../components/ProbesFormGroup";
+import { TolerationFormGroup } from "../components/TolerationFormGroup";
export const ModelStep = ({ version, isEnvironmentDisabled = false, maxAllowedReplica, setMaxAllowedReplica }) => {
const { data, onChangeHandler } = useContext(FormContext);
@@ -65,6 +66,12 @@ export const ModelStep = ({ version, isEnvironmentDisabled = false, maxAllowedRe
onChangeHandler={onChange("image_builder_resource_request")}
errors={get(errors, "image_builder_resource_request")}
/>
+
+
}
/>
diff --git a/ui/src/pages/version/components/forms/steps/TransformerStep.js b/ui/src/pages/version/components/forms/steps/TransformerStep.js
index f2f6fe7ac..40305ba02 100644
--- a/ui/src/pages/version/components/forms/steps/TransformerStep.js
+++ b/ui/src/pages/version/components/forms/steps/TransformerStep.js
@@ -13,6 +13,7 @@ import { LoggerPanel } from "../components/LoggerPanel";
import { ResourcesPanel } from "../components/ResourcesPanel";
import { SelectTransformerPanel } from "../components/SelectTransformerPanel";
import { CPULimitsFormGroup } from "../components/CPULimitsFormGroup";
+import { TolerationFormGroup } from "../components/TolerationFormGroup";
export const TransformerStep = ({ maxAllowedReplica }) => {
const {
@@ -52,6 +53,12 @@ export const TransformerStep = ({ maxAllowedReplica }) => {
onChangeHandler={onChange("transformer.resource_request")}
errors={get(errors, "transformer.resource_request")}
/>
+
+
}
/>
From b0f47e6e296df1ffa3f79fe0c4e3f7437d3dc6a0 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Fri, 26 Jun 2026 20:47:39 +0530
Subject: [PATCH 03/14] update: fixing e2e timeout.
---
scripts/e2e/setup-cluster.sh | 32 ++++++++++++++++++++++++++++----
1 file changed, 28 insertions(+), 4 deletions(-)
diff --git a/scripts/e2e/setup-cluster.sh b/scripts/e2e/setup-cluster.sh
index 883566159..b46c4d0d0 100755
--- a/scripts/e2e/setup-cluster.sh
+++ b/scripts/e2e/setup-cluster.sh
@@ -23,7 +23,7 @@ KNATIVE_NET_ISTIO_VERSION=1.10.1
CERT_MANAGER_VERSION=1.12.2
MINIO_VERSION=3.6.3
KSERVE_VERSION=0.11.0
-TIMEOUT=180s
+TIMEOUT=300s
add_helm_repo() {
@@ -66,9 +66,33 @@ install_istio() {
helm upgrade --install cluster-local-gateway istio/gateway -n istio-system --create-namespace \
-f config/istio/clusterlocal-gateway.yaml --timeout=${TIMEOUT}
- kubectl rollout status deployment/istio-ingressgateway -n istio-system -w --timeout=${TIMEOUT}
- kubectl rollout status deployment/istiod -w -n istio-system --timeout=${TIMEOUT}
- kubectl rollout status deployment/cluster-local-gateway -n istio-system -w --timeout=${TIMEOUT}
+ kubectl rollout status deployment/istio-ingressgateway -n istio-system -w --timeout=${TIMEOUT} || {
+ echo "::group::DEBUG istio-ingressgateway rollout failure"
+ kubectl get pods -n istio-system -o wide
+ kubectl describe pod -n istio-system -l app=istio-ingressgateway
+ kubectl logs -n istio-system -l app=istio-ingressgateway --tail=200 || true
+ kubectl get events -n istio-system --sort-by='.lastTimestamp'
+ echo "::endgroup::"
+ exit 1
+ }
+
+ kubectl rollout status deployment/istiod -w -n istio-system --timeout=${TIMEOUT} || {
+ echo "::group::DEBUG istiod rollout failure"
+ kubectl get pods -n istio-system -o wide
+ kubectl describe pod -n istio-system -l app=istiod
+ kubectl get events -n istio-system --sort-by='.lastTimestamp'
+ echo "::endgroup::"
+ exit 1
+ }
+
+ kubectl rollout status deployment/cluster-local-gateway -n istio-system -w --timeout=${TIMEOUT} || {
+ echo "::group::DEBUG cluster-local-gateway rollout failure"
+ kubectl get pods -n istio-system -o wide
+ kubectl describe pod -n istio-system -l app=cluster-local-gateway
+ kubectl get events -n istio-system --sort-by='.lastTimestamp'
+ echo "::endgroup::"
+ exit 1
+ }
kubectl apply --server-side -f config/istio/ingress-class.yaml
From e264a4618dc210616b6b6295c944a592e6f4ef2b Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Fri, 26 Jun 2026 21:06:10 +0530
Subject: [PATCH 04/14] update: fixing e2e timeout.
---
scripts/e2e/config/istio/ingress-gateway.yaml | 11 +++++++++++
1 file changed, 11 insertions(+)
diff --git a/scripts/e2e/config/istio/ingress-gateway.yaml b/scripts/e2e/config/istio/ingress-gateway.yaml
index c92773f23..2008ffd01 100644
--- a/scripts/e2e/config/istio/ingress-gateway.yaml
+++ b/scripts/e2e/config/istio/ingress-gateway.yaml
@@ -5,3 +5,14 @@ resources:
requests:
cpu: 50m
memory: 64Mi
+
+podSecurityContext:
+ runAsUser: 1337
+ runAsGroup: 1337
+ runAsNonRoot: true
+ fsGroup: 1337
+
+securityContext:
+ runAsUser: 1337
+ runAsGroup: 1337
+ runAsNonRoot: true
\ No newline at end of file
From 988640275180444f364dc8c0e36c388db329bdf1 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Fri, 26 Jun 2026 21:16:31 +0530
Subject: [PATCH 05/14] update: fixing e2e timeout.
---
scripts/e2e/config/istio/clusterlocal-gateway.yaml | 7 ++++++-
scripts/e2e/config/istio/ingress-gateway.yaml | 6 ------
2 files changed, 6 insertions(+), 7 deletions(-)
diff --git a/scripts/e2e/config/istio/clusterlocal-gateway.yaml b/scripts/e2e/config/istio/clusterlocal-gateway.yaml
index 22c6eff65..989313523 100644
--- a/scripts/e2e/config/istio/clusterlocal-gateway.yaml
+++ b/scripts/e2e/config/istio/clusterlocal-gateway.yaml
@@ -15,6 +15,11 @@ resources:
cpu: "1"
memory: 1Gi
+securityContext:
+ runAsUser: 1337
+ runAsGroup: 1337
+ runAsNonRoot: true
+
service:
type: ClusterIP
ports:
@@ -36,4 +41,4 @@ service:
name: http2-prometheus
- port: 15032
targetPort: 15032
- name: http2-tracing
+ name: http2-tracing
\ No newline at end of file
diff --git a/scripts/e2e/config/istio/ingress-gateway.yaml b/scripts/e2e/config/istio/ingress-gateway.yaml
index 2008ffd01..df4c55ebb 100644
--- a/scripts/e2e/config/istio/ingress-gateway.yaml
+++ b/scripts/e2e/config/istio/ingress-gateway.yaml
@@ -6,12 +6,6 @@ resources:
cpu: 50m
memory: 64Mi
-podSecurityContext:
- runAsUser: 1337
- runAsGroup: 1337
- runAsNonRoot: true
- fsGroup: 1337
-
securityContext:
runAsUser: 1337
runAsGroup: 1337
From 2c73da0a7d62947f573f2e8283bb3288ada85f66 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Wed, 8 Jul 2026 03:25:30 +0530
Subject: [PATCH 06/14] fix: fixed merge toleration.
---
api/api/validator.go | 8 ++++
api/api/validator_test.go | 65 +++++++++++++++++++++++++++++++
api/cluster/resource/templater.go | 5 ++-
3 files changed, 77 insertions(+), 1 deletion(-)
diff --git a/api/api/validator.go b/api/api/validator.go
index 9d05f08fb..b5247f8de 100644
--- a/api/api/validator.go
+++ b/api/api/validator.go
@@ -62,6 +62,14 @@ func validateRequest(validators ...requestValidator) error {
func resourceRequestValidation(endpoint *models.VersionEndpoint) requestValidator {
return newFuncValidate(func() error {
+ // Validate transformer tolerations independently: the transformer has its own
+ // resource request and may define tolerations even when the predictor's is nil.
+ if endpoint.Transformer != nil && endpoint.Transformer.ResourceRequest != nil {
+ if err := validateTolerations(endpoint.Transformer.ResourceRequest.Tolerations); err != nil {
+ return fmt.Errorf("invalid toleration in transformer resource request: %w", err)
+ }
+ }
+
if endpoint.ResourceRequest == nil {
return nil
}
diff --git a/api/api/validator_test.go b/api/api/validator_test.go
index 23d1000af..9a940b9cd 100644
--- a/api/api/validator_test.go
+++ b/api/api/validator_test.go
@@ -182,3 +182,68 @@ func TestResourceRequestValidation_Tolerations(t *testing.T) {
})
}
}
+
+func TestResourceRequestValidation_TransformerTolerations(t *testing.T) {
+ tests := []struct {
+ name string
+ endpoint *models.VersionEndpoint
+ wantErr bool
+ errContains string
+ }{
+ {
+ name: "nil transformer passes",
+ endpoint: &models.VersionEndpoint{},
+ },
+ {
+ name: "transformer with nil resource request passes",
+ endpoint: &models.VersionEndpoint{
+ Transformer: &models.Transformer{Enabled: true},
+ },
+ },
+ {
+ name: "transformer with valid tolerations passes",
+ endpoint: &models.VersionEndpoint{
+ Transformer: &models.Transformer{
+ Enabled: true,
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ Tolerations: []corev1.Toleration{
+ {Key: "dedicated", Operator: corev1.TolerationOpEqual, Value: "ml", Effect: corev1.TaintEffectNoSchedule},
+ },
+ },
+ },
+ },
+ },
+ {
+ name: "transformer with invalid toleration fails even when predictor resource request is nil",
+ endpoint: &models.VersionEndpoint{
+ Transformer: &models.Transformer{
+ Enabled: true,
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ Tolerations: []corev1.Toleration{
+ {Key: "k", Operator: corev1.TolerationOpExists, Value: "should-be-empty", Effect: corev1.TaintEffectNoSchedule},
+ },
+ },
+ },
+ },
+ wantErr: true,
+ errContains: "transformer resource request",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ validator := resourceRequestValidation(tt.endpoint)
+ err := validator.validate()
+ if tt.wantErr {
+ assert.Error(t, err)
+ if tt.errContains != "" {
+ assert.Contains(t, err.Error(), tt.errContains)
+ }
+ } else {
+ assert.NoError(t, err)
+ }
+ })
+ }
+}
diff --git a/api/cluster/resource/templater.go b/api/cluster/resource/templater.go
index 06864f559..73aff3284 100644
--- a/api/cluster/resource/templater.go
+++ b/api/cluster/resource/templater.go
@@ -231,7 +231,10 @@ func (t *InferenceServiceTemplater) createPredictorSpec(modelService *models.Ser
resources.Limits[resourceType] = resourceQuantity
nodeSelector = gpuConfig.NodeSelector
- tolerations = gpuConfig.Tolerations
+ // Copy into a slice we own rather than aliasing the shared
+ // deploymentConfig backing array (a subsequent append must not
+ // mutate the config or leak across deployments).
+ tolerations = append(tolerations, gpuConfig.Tolerations...)
}
}
}
From e65bf29f83cde2767709ea57d9f23d55fc8e4852 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Wed, 8 Jul 2026 03:43:26 +0530
Subject: [PATCH 07/14] removed unwanted test
---
api/api/validator.go | 36 ++++++++++++++++++++----------------
api/api/validator_test.go | 29 ++++++++++++++++++++++++++++-
2 files changed, 48 insertions(+), 17 deletions(-)
diff --git a/api/api/validator.go b/api/api/validator.go
index b5247f8de..ddd1a4e6f 100644
--- a/api/api/validator.go
+++ b/api/api/validator.go
@@ -62,28 +62,32 @@ func validateRequest(validators ...requestValidator) error {
func resourceRequestValidation(endpoint *models.VersionEndpoint) requestValidator {
return newFuncValidate(func() error {
- // Validate transformer tolerations independently: the transformer has its own
- // resource request and may define tolerations even when the predictor's is nil.
- if endpoint.Transformer != nil && endpoint.Transformer.ResourceRequest != nil {
- if err := validateTolerations(endpoint.Transformer.ResourceRequest.Tolerations); err != nil {
- return fmt.Errorf("invalid toleration in transformer resource request: %w", err)
+ if endpoint.ResourceRequest != nil {
+ if endpoint.ResourceRequest.MinReplica > endpoint.ResourceRequest.MaxReplica {
+ return fmt.Errorf("min replica must be less or equal to max replica")
}
- }
- if endpoint.ResourceRequest == nil {
- return nil
- }
+ if endpoint.ResourceRequest.MaxReplica < 1 {
+ return fmt.Errorf("max replica must be greater than 0")
+ }
- if endpoint.ResourceRequest.MinReplica > endpoint.ResourceRequest.MaxReplica {
- return fmt.Errorf("min replica must be less or equal to max replica")
+ if err := validateTolerations(endpoint.ResourceRequest.Tolerations); err != nil {
+ return fmt.Errorf("invalid toleration in resource request: %w", err)
+ }
}
- if endpoint.ResourceRequest.MaxReplica < 1 {
- return fmt.Errorf("max replica must be greater than 0")
- }
+ if endpoint.Transformer != nil && endpoint.Transformer.ResourceRequest != nil {
+ if endpoint.Transformer.ResourceRequest.MinReplica > endpoint.Transformer.ResourceRequest.MaxReplica {
+ return fmt.Errorf("transformer min replica must be less or equal to max replica")
+ }
+
+ if endpoint.Transformer.ResourceRequest.MaxReplica < 1 {
+ return fmt.Errorf("transformer max replica must be greater than 0")
+ }
- if err := validateTolerations(endpoint.ResourceRequest.Tolerations); err != nil {
- return fmt.Errorf("invalid toleration in resource request: %w", err)
+ if err := validateTolerations(endpoint.Transformer.ResourceRequest.Tolerations); err != nil {
+ return fmt.Errorf("invalid toleration in transformer resource request: %w", err)
+ }
}
return nil
diff --git a/api/api/validator_test.go b/api/api/validator_test.go
index 9a940b9cd..65085e2d4 100644
--- a/api/api/validator_test.go
+++ b/api/api/validator_test.go
@@ -142,6 +142,7 @@ func TestResourceRequestValidation_Tolerations(t *testing.T) {
tests := []struct {
name string
req *models.ResourceRequest
+ trans *models.Transformer
wantErr bool
}{
{
@@ -157,6 +158,17 @@ func TestResourceRequestValidation_Tolerations(t *testing.T) {
},
},
},
+ {
+ name: "valid transformer resource request with tolerations passes",
+ trans: &models.Transformer{
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ Tolerations: []corev1.Toleration{
+ {Key: "dedicated", Operator: corev1.TolerationOpEqual, Value: "transformer", Effect: corev1.TaintEffectNoSchedule},
+ },
+ },
+ },
+ },
{
name: "resource request with invalid toleration operator fails",
req: &models.ResourceRequest{
@@ -167,11 +179,26 @@ func TestResourceRequestValidation_Tolerations(t *testing.T) {
},
wantErr: true,
},
+ {
+ name: "transformer resource request with invalid toleration operator fails",
+ trans: &models.Transformer{
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ Tolerations: []corev1.Toleration{
+ {Key: "k", Operator: "BadOp", Value: "v"},
+ },
+ },
+ },
+ wantErr: true,
+ },
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- endpoint := &models.VersionEndpoint{ResourceRequest: tt.req}
+ endpoint := &models.VersionEndpoint{
+ ResourceRequest: tt.req,
+ Transformer: tt.trans,
+ }
validator := resourceRequestValidation(endpoint)
err := validator.validate()
if tt.wantErr {
From 8e94650ca3272209fc099c5bcf95db0c8666d8d8 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Wed, 8 Jul 2026 03:43:26 +0530
Subject: [PATCH 08/14] fix: reconcile transformer resource request validation
Validate transformer min/max replica and tolerations alongside the
predictor's, without short-circuiting when the predictor resource
request is nil.
Co-Authored-By: Claude Opus 4.8
---
api/api/validator.go | 36 ++++++++++++++++++++----------------
api/api/validator_test.go | 29 ++++++++++++++++++++++++++++-
2 files changed, 48 insertions(+), 17 deletions(-)
diff --git a/api/api/validator.go b/api/api/validator.go
index b5247f8de..ddd1a4e6f 100644
--- a/api/api/validator.go
+++ b/api/api/validator.go
@@ -62,28 +62,32 @@ func validateRequest(validators ...requestValidator) error {
func resourceRequestValidation(endpoint *models.VersionEndpoint) requestValidator {
return newFuncValidate(func() error {
- // Validate transformer tolerations independently: the transformer has its own
- // resource request and may define tolerations even when the predictor's is nil.
- if endpoint.Transformer != nil && endpoint.Transformer.ResourceRequest != nil {
- if err := validateTolerations(endpoint.Transformer.ResourceRequest.Tolerations); err != nil {
- return fmt.Errorf("invalid toleration in transformer resource request: %w", err)
+ if endpoint.ResourceRequest != nil {
+ if endpoint.ResourceRequest.MinReplica > endpoint.ResourceRequest.MaxReplica {
+ return fmt.Errorf("min replica must be less or equal to max replica")
}
- }
- if endpoint.ResourceRequest == nil {
- return nil
- }
+ if endpoint.ResourceRequest.MaxReplica < 1 {
+ return fmt.Errorf("max replica must be greater than 0")
+ }
- if endpoint.ResourceRequest.MinReplica > endpoint.ResourceRequest.MaxReplica {
- return fmt.Errorf("min replica must be less or equal to max replica")
+ if err := validateTolerations(endpoint.ResourceRequest.Tolerations); err != nil {
+ return fmt.Errorf("invalid toleration in resource request: %w", err)
+ }
}
- if endpoint.ResourceRequest.MaxReplica < 1 {
- return fmt.Errorf("max replica must be greater than 0")
- }
+ if endpoint.Transformer != nil && endpoint.Transformer.ResourceRequest != nil {
+ if endpoint.Transformer.ResourceRequest.MinReplica > endpoint.Transformer.ResourceRequest.MaxReplica {
+ return fmt.Errorf("transformer min replica must be less or equal to max replica")
+ }
+
+ if endpoint.Transformer.ResourceRequest.MaxReplica < 1 {
+ return fmt.Errorf("transformer max replica must be greater than 0")
+ }
- if err := validateTolerations(endpoint.ResourceRequest.Tolerations); err != nil {
- return fmt.Errorf("invalid toleration in resource request: %w", err)
+ if err := validateTolerations(endpoint.Transformer.ResourceRequest.Tolerations); err != nil {
+ return fmt.Errorf("invalid toleration in transformer resource request: %w", err)
+ }
}
return nil
diff --git a/api/api/validator_test.go b/api/api/validator_test.go
index 9a940b9cd..65085e2d4 100644
--- a/api/api/validator_test.go
+++ b/api/api/validator_test.go
@@ -142,6 +142,7 @@ func TestResourceRequestValidation_Tolerations(t *testing.T) {
tests := []struct {
name string
req *models.ResourceRequest
+ trans *models.Transformer
wantErr bool
}{
{
@@ -157,6 +158,17 @@ func TestResourceRequestValidation_Tolerations(t *testing.T) {
},
},
},
+ {
+ name: "valid transformer resource request with tolerations passes",
+ trans: &models.Transformer{
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ Tolerations: []corev1.Toleration{
+ {Key: "dedicated", Operator: corev1.TolerationOpEqual, Value: "transformer", Effect: corev1.TaintEffectNoSchedule},
+ },
+ },
+ },
+ },
{
name: "resource request with invalid toleration operator fails",
req: &models.ResourceRequest{
@@ -167,11 +179,26 @@ func TestResourceRequestValidation_Tolerations(t *testing.T) {
},
wantErr: true,
},
+ {
+ name: "transformer resource request with invalid toleration operator fails",
+ trans: &models.Transformer{
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ Tolerations: []corev1.Toleration{
+ {Key: "k", Operator: "BadOp", Value: "v"},
+ },
+ },
+ },
+ wantErr: true,
+ },
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- endpoint := &models.VersionEndpoint{ResourceRequest: tt.req}
+ endpoint := &models.VersionEndpoint{
+ ResourceRequest: tt.req,
+ Transformer: tt.trans,
+ }
validator := resourceRequestValidation(endpoint)
err := validator.validate()
if tt.wantErr {
From 6b963e34781d047d781792e7c8ad4daccecf1299 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Wed, 8 Jul 2026 03:52:29 +0530
Subject: [PATCH 09/14] revert: remove orphaned UTF-8 sanitize test
Completes revert 345a1d5, which reverted the sanitizeEndpoint UTF-8 fix
(b0ffd2f) but left its test file behind, causing TestSanitizeEndpoint_*
to fail against the reverted code.
Co-Authored-By: Claude Opus 4.8
---
.../version_endpoint_storage_unit_test.go | 33 -------------------
1 file changed, 33 deletions(-)
delete mode 100644 api/storage/version_endpoint_storage_unit_test.go
diff --git a/api/storage/version_endpoint_storage_unit_test.go b/api/storage/version_endpoint_storage_unit_test.go
deleted file mode 100644
index d4cc42fd1..000000000
--- a/api/storage/version_endpoint_storage_unit_test.go
+++ /dev/null
@@ -1,33 +0,0 @@
-package storage
-
-import (
- "strings"
- "testing"
- "unicode/utf8"
-
- "github.com/caraml-dev/merlin/models"
- "github.com/stretchr/testify/assert"
-)
-
-func TestSanitizeEndpoint_RemovesInvalidUTF8(t *testing.T) {
- endpoint := &models.VersionEndpoint{
- Message: "prefix" + string([]byte{0xe2, 0x94}) + "suffix",
- }
-
- sanitizeEndpoint(endpoint)
-
- assert.Equal(t, "prefixsuffix", endpoint.Message)
- assert.True(t, utf8.ValidString(endpoint.Message))
-}
-
-func TestSanitizeEndpoint_TruncatesToValidUTF8(t *testing.T) {
- message := strings.Repeat("a", maxMessageChar-1) + "e" + "\u0301"
- endpoint := &models.VersionEndpoint{
- Message: message,
- }
-
- sanitizeEndpoint(endpoint)
-
- assert.Equal(t, maxMessageChar, len(endpoint.Message))
- assert.True(t, utf8.ValidString(endpoint.Message))
-}
From 09c28c5c80e2de8124ffbbc9ec10e3c1a9876204 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Wed, 8 Jul 2026 04:42:09 +0530
Subject: [PATCH 10/14] fix: ui fixed
---
.github/workflows/api-image-fast.yml | 178 ++++++++++++++++++
ui/src/services/transformer/Transformer.js | 7 +-
.../version_endpoint/VersionEndpoint.js | 5 +
3 files changed, 189 insertions(+), 1 deletion(-)
create mode 100644 .github/workflows/api-image-fast.yml
diff --git a/.github/workflows/api-image-fast.yml b/.github/workflows/api-image-fast.yml
new file mode 100644
index 000000000..9d65b71e0
--- /dev/null
+++ b/.github/workflows/api-image-fast.yml
@@ -0,0 +1,178 @@
+name: Fast API Image CI
+
+on:
+ workflow_dispatch:
+ inputs:
+ image_tag:
+ description: Docker image tag to use. Defaults to the commit SHA.
+ required: false
+ type: string
+ push_image:
+ description: Push the built image to GHCR.
+ required: true
+ default: true
+ type: boolean
+
+env:
+ GO_VERSION: "1.22"
+ DOCKER_REGISTRY: ghcr.io
+
+jobs:
+ prepare:
+ runs-on: ubuntu-latest
+ outputs:
+ image_tag: ${{ steps.meta.outputs.image_tag }}
+ artifact_name: ${{ steps.meta.outputs.artifact_name }}
+ transformer_artifact_name: ${{ steps.meta.outputs.transformer_artifact_name }}
+ steps:
+ - id: meta
+ env:
+ INPUT_IMAGE_TAG: ${{ inputs.image_tag }}
+ run: |
+ IMAGE_TAG="${INPUT_IMAGE_TAG}"
+ if [ -z "${IMAGE_TAG}" ]; then
+ IMAGE_TAG="${GITHUB_SHA}"
+ fi
+
+ echo "image_tag=${IMAGE_TAG}" >> "$GITHUB_OUTPUT"
+ echo "artifact_name=merlin.${IMAGE_TAG}.tar" >> "$GITHUB_OUTPUT"
+ echo "transformer_artifact_name=merlin-transformer.${IMAGE_TAG}.tar" >> "$GITHUB_OUTPUT"
+
+ test-api:
+ runs-on: ubuntu-latest
+ services:
+ postgres:
+ image: postgres:12.4
+ env:
+ POSTGRES_DB: postgres
+ POSTGRES_USER: postgres
+ POSTGRES_PASSWORD: postgres
+ ports:
+ - 5432:5432
+ steps:
+ - uses: actions/checkout@v4
+ - uses: actions/setup-go@v5
+ with:
+ go-version: ${{ env.GO_VERSION }}
+ cache-dependency-path: api/go.sum
+ - name: Install dependencies
+ run: |
+ make setup
+ make init-dep-api
+ - name: Test API files
+ env:
+ POSTGRES_HOST: localhost
+ POSTGRES_DB: postgres
+ POSTGRES_USER: postgres
+ POSTGRES_PASSWORD: postgres
+ run: make it-test-api-ci
+
+ build-ui:
+ runs-on: ubuntu-latest
+ steps:
+ - uses: actions/checkout@v4
+ - uses: actions/setup-node@v4
+ with:
+ node-version: 20
+ cache: yarn
+ cache-dependency-path: ui/yarn.lock
+ - name: Install dependencies
+ run: make init-dep-ui
+ - name: Build UI static files
+ run: make build-ui
+ - name: Publish UI artifact
+ uses: actions/upload-artifact@v4
+ with:
+ name: merlin-ui-dist-fast
+ path: ui/build/
+
+ build-transformer:
+ runs-on: ubuntu-latest
+ needs:
+ - prepare
+ steps:
+ - uses: actions/checkout@v4
+ - uses: actions/setup-go@v5
+ with:
+ go-version: ${{ env.GO_VERSION }}
+ cache-dependency-path: api/go.sum
+ - name: Install dependencies
+ run: make init-dep-api
+ - name: Build Standard Transformer
+ run: make build-transformer
+ - name: Build Standard Transformer Docker image
+ run: docker build -t merlin-transformer:${{ needs.prepare.outputs.image_tag }} -f transformer.Dockerfile .
+ - name: Save Standard Transformer Docker image
+ run: docker image save --output "${{ needs.prepare.outputs.transformer_artifact_name }}" "merlin-transformer:${{ needs.prepare.outputs.image_tag }}"
+ - name: Publish Standard Transformer Docker artifact
+ uses: actions/upload-artifact@v4
+ with:
+ name: ${{ needs.prepare.outputs.transformer_artifact_name }}
+ path: ${{ needs.prepare.outputs.transformer_artifact_name }}
+
+ build-api-image:
+ runs-on: ubuntu-latest
+ needs:
+ - prepare
+ - test-api
+ - build-ui
+ steps:
+ - uses: actions/checkout@v4
+ - name: Download UI artifact
+ uses: actions/download-artifact@v4
+ with:
+ name: merlin-ui-dist-fast
+ path: ui/build
+ - name: Build API Docker image
+ run: docker build -t merlin:${{ needs.prepare.outputs.image_tag }} -f Dockerfile .
+ - name: Save API Docker image
+ run: docker image save --output "${{ needs.prepare.outputs.artifact_name }}" "merlin:${{ needs.prepare.outputs.image_tag }}"
+ - name: Publish API Docker artifact
+ uses: actions/upload-artifact@v4
+ with:
+ name: ${{ needs.prepare.outputs.artifact_name }}
+ path: ${{ needs.prepare.outputs.artifact_name }}
+
+ push-api-image:
+ if: ${{ inputs.push_image }}
+ runs-on: ubuntu-latest
+ needs:
+ - prepare
+ - build-api-image
+ permissions:
+ packages: write
+ steps:
+ - name: Download API Docker artifact
+ uses: actions/download-artifact@v4
+ with:
+ name: ${{ needs.prepare.outputs.artifact_name }}
+ - name: Push Docker image
+ env:
+ IMAGE_TAG: ghcr.io/${{ github.repository }}/merlin:${{ needs.prepare.outputs.image_tag }}
+ run: |
+ docker login ${{ env.DOCKER_REGISTRY }} -u ${{ github.actor }} -p ${{ secrets.GITHUB_TOKEN }}
+ docker image load --input "${{ needs.prepare.outputs.artifact_name }}"
+ docker tag "merlin:${{ needs.prepare.outputs.image_tag }}" "${IMAGE_TAG}"
+ docker push "${IMAGE_TAG}"
+
+ push-transformer-image:
+ if: ${{ inputs.push_image }}
+ runs-on: ubuntu-latest
+ needs:
+ - prepare
+ - build-transformer
+ permissions:
+ packages: write
+ steps:
+ - name: Download Standard Transformer Docker artifact
+ uses: actions/download-artifact@v4
+ with:
+ name: ${{ needs.prepare.outputs.transformer_artifact_name }}
+ - name: Push Docker image
+ env:
+ IMAGE_TAG: ghcr.io/${{ github.repository }}/merlin-transformer:${{ needs.prepare.outputs.image_tag }}
+ run: |
+ docker login ${{ env.DOCKER_REGISTRY }} -u ${{ github.actor }} -p ${{ secrets.GITHUB_TOKEN }}
+ docker image load --input "${{ needs.prepare.outputs.transformer_artifact_name }}"
+ docker tag "merlin-transformer:${{ needs.prepare.outputs.image_tag }}" "${IMAGE_TAG}"
+ docker push "${IMAGE_TAG}"
diff --git a/ui/src/services/transformer/Transformer.js b/ui/src/services/transformer/Transformer.js
index 01bd4fb5b..b827f4bd0 100644
--- a/ui/src/services/transformer/Transformer.js
+++ b/ui/src/services/transformer/Transformer.js
@@ -23,7 +23,8 @@ export class Transformer {
max_replica: process.env.REACT_APP_ENVIRONMENT === "production" ? 4 : 2,
cpu_request: "500m",
cpu_limit: "",
- memory_request: "512Mi"
+ memory_request: "512Mi",
+ tolerations: [],
};
this.env_vars = [];
@@ -60,6 +61,10 @@ export class Transformer {
transformer.secrets = [];
}
+ if (transformer.resource_request && !transformer.resource_request.tolerations) {
+ transformer.resource_request.tolerations = [];
+ }
+
return transformer;
}
diff --git a/ui/src/services/version_endpoint/VersionEndpoint.js b/ui/src/services/version_endpoint/VersionEndpoint.js
index 465058ddc..40c792ec1 100644
--- a/ui/src/services/version_endpoint/VersionEndpoint.js
+++ b/ui/src/services/version_endpoint/VersionEndpoint.js
@@ -30,6 +30,7 @@ export class VersionEndpoint {
memory_request: "512Mi",
liveness_probe: null,
readiness_probe: null,
+ tolerations: [],
};
this.image_builder_resource_request = {
@@ -74,6 +75,10 @@ export class VersionEndpoint {
}
}
+ if (versionEndpoint.resource_request && !versionEndpoint.resource_request.tolerations) {
+ versionEndpoint.resource_request.tolerations = [];
+ }
+
if (json.transformer) {
versionEndpoint.transformer = Transformer.fromJson(json.transformer);
}
From a59b33557f36bbbbfc0677b1b6a2c990c2195b12 Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Mon, 20 Jul 2026 04:54:07 +0530
Subject: [PATCH 11/14] added get nodes in merlin
---
api/api/node_pool_api.go | 50 ++++++++
api/api/router.go | 3 +
api/api/validator.go | 18 +++
api/client/model_resource_request.go | 51 ++++++--
api/cluster/controller.go | 6 +
api/cluster/mocks/controller.go | 30 +++++
api/cluster/resource/templater.go | 20 +++-
api/cluster/resource/templater_test.go | 112 ++++++++++++++++++
api/cmd/api/main.go | 2 +
api/models/node_pool.go | 106 +++++++++++++++++
api/models/node_pool_test.go | 68 +++++++++++
api/models/resource_request.go | 2 +
api/service/node_pool_service.go | 53 +++++++++
swagger.yaml | 45 +++++++
ui/src/components/ResourcesConfigTable.js | 11 ++
.../forms/DeployModelVersionForm.js | 7 ++
.../forms/components/NodePoolSelect.js | 83 +++++++++++++
.../forms/components/NodeSelectorFormGroup.js | 100 ++++++++++++++++
.../components/forms/steps/ModelStep.js | 25 +++-
.../components/forms/steps/TransformerStep.js | 6 +
ui/src/services/transformer/Transformer.js | 5 +
.../version_endpoint/VersionEndpoint.js | 5 +
22 files changed, 795 insertions(+), 13 deletions(-)
create mode 100644 api/api/node_pool_api.go
create mode 100644 api/models/node_pool.go
create mode 100644 api/models/node_pool_test.go
create mode 100644 api/service/node_pool_service.go
create mode 100644 ui/src/pages/version/components/forms/components/NodePoolSelect.js
create mode 100644 ui/src/pages/version/components/forms/components/NodeSelectorFormGroup.js
diff --git a/api/api/node_pool_api.go b/api/api/node_pool_api.go
new file mode 100644
index 000000000..d355d18b2
--- /dev/null
+++ b/api/api/node_pool_api.go
@@ -0,0 +1,50 @@
+// Copyright 2020 The Merlin Authors
+//
+// 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 api
+
+import (
+ "errors"
+ "fmt"
+ "net/http"
+
+ "gorm.io/gorm"
+)
+
+// NodePoolController serves node pool discovery for a deployment environment.
+type NodePoolController struct {
+ *AppContext
+}
+
+// ListNodePools returns the schedulable node pools (taint + shared labels) of the
+// cluster backing the given environment, so the UI can offer them for placement.
+func (c *NodePoolController) ListNodePools(r *http.Request, vars map[string]string, _ interface{}) *Response {
+ ctx := r.Context()
+
+ environmentName := vars["environment_name"]
+ env, err := c.EnvironmentService.GetEnvironment(environmentName)
+ if err != nil {
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return NotFound(fmt.Sprintf("Environment not found: %v", err))
+ }
+ return InternalServerError(fmt.Sprintf("Error getting environment: %v", err))
+ }
+
+ nodePools, err := c.NodePoolService.ListNodePools(ctx, env.Cluster)
+ if err != nil {
+ return InternalServerError(fmt.Sprintf("Error listing node pools: %v", err))
+ }
+
+ return Ok(nodePools)
+}
diff --git a/api/api/router.go b/api/api/router.go
index 6fc8902d8..5b584bf22 100644
--- a/api/api/router.go
+++ b/api/api/router.go
@@ -62,6 +62,7 @@ type AppContext struct {
VersionImageService service.VersionImageService
EndpointsService service.EndpointsService
LogService service.LogService
+ NodePoolService service.NodePoolService
PredictionJobService service.PredictionJobService
SecretService service.SecretService
ModelEndpointAlertService service.ModelEndpointAlertService
@@ -168,6 +169,7 @@ func NewRouter(appCtx AppContext) (*mux.Router, error) {
endpointsController := EndpointsController{&appCtx}
predictionJobController := PredictionJobController{&appCtx}
logController := LogController{&appCtx}
+ nodePoolController := NodePoolController{&appCtx}
secretController := SecretsController{&appCtx}
alertsController := AlertsController{&appCtx}
transformerController := TransformerController{&appCtx}
@@ -176,6 +178,7 @@ func NewRouter(appCtx AppContext) (*mux.Router, error) {
routes := []Route{
// Environment API
{http.MethodGet, "/environments", nil, environmentController.ListEnvironments, "ListEnvironments"},
+ {http.MethodGet, "/environments/{environment_name}/node-pools", nil, nodePoolController.ListNodePools, "ListNodePools"},
// Project API
{http.MethodGet, "/projects/{project_id:[0-9]+}", nil, projectsController.GetProject, "GetProject"},
diff --git a/api/api/validator.go b/api/api/validator.go
index ddd1a4e6f..ea361dedd 100644
--- a/api/api/validator.go
+++ b/api/api/validator.go
@@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
+ "strings"
"golang.org/x/exp/slices"
corev1 "k8s.io/api/core/v1"
@@ -74,6 +75,10 @@ func resourceRequestValidation(endpoint *models.VersionEndpoint) requestValidato
if err := validateTolerations(endpoint.ResourceRequest.Tolerations); err != nil {
return fmt.Errorf("invalid toleration in resource request: %w", err)
}
+
+ if err := validateNodeSelector(endpoint.ResourceRequest.NodeSelector); err != nil {
+ return fmt.Errorf("invalid node selector in resource request: %w", err)
+ }
}
if endpoint.Transformer != nil && endpoint.Transformer.ResourceRequest != nil {
@@ -88,6 +93,10 @@ func resourceRequestValidation(endpoint *models.VersionEndpoint) requestValidato
if err := validateTolerations(endpoint.Transformer.ResourceRequest.Tolerations); err != nil {
return fmt.Errorf("invalid toleration in transformer resource request: %w", err)
}
+
+ if err := validateNodeSelector(endpoint.Transformer.ResourceRequest.NodeSelector); err != nil {
+ return fmt.Errorf("invalid node selector in transformer resource request: %w", err)
+ }
}
return nil
@@ -125,6 +134,15 @@ func validateTolerations(tolerations []corev1.Toleration) error {
return nil
}
+func validateNodeSelector(nodeSelector map[string]string) error {
+ for k := range nodeSelector {
+ if strings.TrimSpace(k) == "" {
+ return fmt.Errorf("node selector must not contain an empty label key")
+ }
+ }
+ return nil
+}
+
func customModelValidation(model *models.Model, version *models.Version) requestValidator {
return newFuncValidate(func() error {
if model.Type == models.ModelTypeCustom {
diff --git a/api/client/model_resource_request.go b/api/client/model_resource_request.go
index bd719c229..de104ca8f 100644
--- a/api/client/model_resource_request.go
+++ b/api/client/model_resource_request.go
@@ -19,14 +19,15 @@ var _ MappedNullable = &ResourceRequest{}
// ResourceRequest struct for ResourceRequest
type ResourceRequest struct {
- MinReplica *int32 `json:"min_replica,omitempty"`
- MaxReplica *int32 `json:"max_replica,omitempty"`
- CpuRequest *string `json:"cpu_request,omitempty"`
- CpuLimit *string `json:"cpu_limit,omitempty"`
- MemoryRequest *string `json:"memory_request,omitempty"`
- GpuName *string `json:"gpu_name,omitempty"`
- GpuRequest *string `json:"gpu_request,omitempty"`
- Tolerations []Toleration `json:"tolerations,omitempty"`
+ MinReplica *int32 `json:"min_replica,omitempty"`
+ MaxReplica *int32 `json:"max_replica,omitempty"`
+ CpuRequest *string `json:"cpu_request,omitempty"`
+ CpuLimit *string `json:"cpu_limit,omitempty"`
+ MemoryRequest *string `json:"memory_request,omitempty"`
+ GpuName *string `json:"gpu_name,omitempty"`
+ GpuRequest *string `json:"gpu_request,omitempty"`
+ Tolerations []Toleration `json:"tolerations,omitempty"`
+ NodeSelector *map[string]string `json:"node_selector,omitempty"`
}
// NewResourceRequest instantiates a new ResourceRequest object
@@ -301,6 +302,37 @@ func (o *ResourceRequest) SetTolerations(v []Toleration) {
o.Tolerations = v
}
+// GetNodeSelector returns the NodeSelector field value if set, zero value otherwise.
+func (o *ResourceRequest) GetNodeSelector() map[string]string {
+ if o == nil || IsNil(o.NodeSelector) {
+ var ret map[string]string
+ return ret
+ }
+ return *o.NodeSelector
+}
+
+// GetNodeSelectorOk returns a tuple with the NodeSelector field value if set, nil otherwise
+// and a boolean to check if the value has been set.
+func (o *ResourceRequest) GetNodeSelectorOk() (*map[string]string, bool) {
+ if o == nil || IsNil(o.NodeSelector) {
+ return nil, false
+ }
+ return o.NodeSelector, true
+}
+
+// HasNodeSelector returns a boolean if a field has been set.
+func (o *ResourceRequest) HasNodeSelector() bool {
+ if o != nil && !IsNil(o.NodeSelector) {
+ return true
+ }
+ return false
+}
+
+// SetNodeSelector gets a reference to the given map[string]string and assigns it to the NodeSelector field.
+func (o *ResourceRequest) SetNodeSelector(v map[string]string) {
+ o.NodeSelector = &v
+}
+
func (o ResourceRequest) MarshalJSON() ([]byte, error) {
toSerialize, err := o.ToMap()
if err != nil {
@@ -335,6 +367,9 @@ func (o ResourceRequest) ToMap() (map[string]interface{}, error) {
if !IsNil(o.Tolerations) {
toSerialize["tolerations"] = o.Tolerations
}
+ if !IsNil(o.NodeSelector) {
+ toSerialize["node_selector"] = o.NodeSelector
+ }
return toSerialize, nil
}
diff --git a/api/cluster/controller.go b/api/cluster/controller.go
index d89d83f6e..fce2de496 100644
--- a/api/cluster/controller.go
+++ b/api/cluster/controller.go
@@ -55,6 +55,8 @@ type Controller interface {
ListPods(ctx context.Context, namespace, labelSelector string) (*corev1.PodList, error)
StreamPodLogs(ctx context.Context, namespace, podName string, opts *corev1.PodLogOptions) (io.ReadCloser, error)
+ ListNodes(ctx context.Context) (*corev1.NodeList, error)
+
ListJobs(ctx context.Context, namespace, labelSelector string) (*batchv1.JobList, error)
DeleteJob(ctx context.Context, namespace, jobName string, deleteOptions metav1.DeleteOptions) error
DeleteJobs(ctx context.Context, namespace string, deleteOptions metav1.DeleteOptions, listOptions metav1.ListOptions) error
@@ -503,6 +505,10 @@ func (c *controller) StreamPodLogs(ctx context.Context, namespace, podName strin
return c.clusterClient.Pods(namespace).GetLogs(podName, opts).Stream(ctx)
}
+func (c *controller) ListNodes(ctx context.Context) (*corev1.NodeList, error) {
+ return c.clusterClient.Nodes().List(ctx, metav1.ListOptions{})
+}
+
func (c *controller) GetCurrentDeploymentScale(
ctx context.Context,
namespace string,
diff --git a/api/cluster/mocks/controller.go b/api/cluster/mocks/controller.go
index 184bab903..0696fbc9c 100644
--- a/api/cluster/mocks/controller.go
+++ b/api/cluster/mocks/controller.go
@@ -202,6 +202,36 @@ func (_m *Controller) ListJobs(ctx context.Context, namespace string, labelSelec
}
// ListPods provides a mock function with given fields: ctx, namespace, labelSelector
+// ListNodes provides a mock function with given fields: ctx
+func (_m *Controller) ListNodes(ctx context.Context) (*corev1.NodeList, error) {
+ ret := _m.Called(ctx)
+
+ if len(ret) == 0 {
+ panic("no return value specified for ListNodes")
+ }
+
+ var r0 *corev1.NodeList
+ var r1 error
+ if rf, ok := ret.Get(0).(func(context.Context) (*corev1.NodeList, error)); ok {
+ return rf(ctx)
+ }
+ if rf, ok := ret.Get(0).(func(context.Context) *corev1.NodeList); ok {
+ r0 = rf(ctx)
+ } else {
+ if ret.Get(0) != nil {
+ r0 = ret.Get(0).(*corev1.NodeList)
+ }
+ }
+
+ if rf, ok := ret.Get(1).(func(context.Context) error); ok {
+ r1 = rf(ctx)
+ } else {
+ r1 = ret.Error(1)
+ }
+
+ return r0, r1
+}
+
func (_m *Controller) ListPods(ctx context.Context, namespace string, labelSelector string) (*corev1.PodList, error) {
ret := _m.Called(ctx, namespace, labelSelector)
diff --git a/api/cluster/resource/templater.go b/api/cluster/resource/templater.go
index 73aff3284..2b3336242 100644
--- a/api/cluster/resource/templater.go
+++ b/api/cluster/resource/templater.go
@@ -230,10 +230,12 @@ func (t *InferenceServiceTemplater) createPredictorSpec(modelService *models.Ser
resources.Requests[resourceType] = resourceQuantity
resources.Limits[resourceType] = resourceQuantity
- nodeSelector = gpuConfig.NodeSelector
- // Copy into a slice we own rather than aliasing the shared
- // deploymentConfig backing array (a subsequent append must not
- // mutate the config or leak across deployments).
+ // Copy into maps/slices we own rather than aliasing the shared
+ // deploymentConfig (a subsequent write must not mutate the config
+ // or leak across deployments).
+ for k, v := range gpuConfig.NodeSelector {
+ nodeSelector[k] = v
+ }
tolerations = append(tolerations, gpuConfig.Tolerations...)
}
}
@@ -245,6 +247,11 @@ func (t *InferenceServiceTemplater) createPredictorSpec(modelService *models.Ser
tolerations = append(tolerations, modelService.ResourceRequest.Tolerations...)
}
+ // Overlay user-defined node selectors (merged on top of any GPU-derived selector)
+ for k, v := range modelService.ResourceRequest.NodeSelector {
+ nodeSelector[k] = v
+ }
+
// Get user-configured probe settings
var userLivenessConfig *models.ProbeConfig
var userReadinessConfig *models.ProbeConfig
@@ -481,6 +488,11 @@ func (t *InferenceServiceTemplater) createTransformerSpec(
transformerSpec.PodSpec.Tolerations = transformer.ResourceRequest.Tolerations
}
+ // Apply user-defined node selector for transformer pods
+ if len(transformer.ResourceRequest.NodeSelector) > 0 {
+ transformerSpec.PodSpec.NodeSelector = transformer.ResourceRequest.NodeSelector
+ }
+
return transformerSpec, nil
}
diff --git a/api/cluster/resource/templater_test.go b/api/cluster/resource/templater_test.go
index 91ac5afb5..e608f0c2d 100644
--- a/api/cluster/resource/templater_test.go
+++ b/api/cluster/resource/templater_test.go
@@ -5007,3 +5007,115 @@ func TestCreateInferenceServiceSpecWithTolerations(t *testing.T) {
})
}
}
+
+func TestCreateInferenceServiceSpecWithNodeSelector(t *testing.T) {
+ err := labeller.InitKubernetesLabeller("gojek.com/", "caraml.dev/", testEnvironmentName)
+ assert.NoError(t, err)
+ defer func() { _ = labeller.InitKubernetesLabeller("", "", "") }()
+
+ project := mlp.Project{Name: "project"}
+ userNodeSelector := map[string]string{"pool": "workload-optimized"}
+ gpuNodeSelector := map[string]string{"cloud.google.com/gke-accelerator": "nvidia-tesla-t4"}
+
+ baseMeta := models.Metadata{App: "model", Component: models.ComponentModelVersion, Stream: "dsp", Team: "dsp"}
+
+ baseDeployConfig := &config.DeploymentConfig{
+ DefaultModelResourceRequests: defaultModelResourceRequests,
+ DefaultTransformerResourceRequests: defaultTransformerResourceRequests,
+ QueueResourcePercentage: "2",
+ StandardTransformer: standardTransformerConfig,
+ UserContainerCPUDefaultLimit: userContainerCPUDefaultLimit,
+ UserContainerCPULimitRequestFactor: userContainerCPULimitRequestFactor,
+ UserContainerMemoryLimitRequestFactor: userContainerMemoryLimitRequestFactor,
+ DefaultEnvVarsWithoutCPULimits: []corev1.EnvVar{defaultEnvVarWithoutCPULimits},
+ }
+
+ gpuConfig := config.GPUConfig{
+ Name: "NVIDIA T4",
+ Values: []string{"1"},
+ ResourceType: "nvidia.com/gpu",
+ NodeSelector: gpuNodeSelector,
+ }
+
+ tests := []struct {
+ name string
+ modelSvc *models.Service
+ deployConfig *config.DeploymentConfig
+ checkFn func(t *testing.T, infSvc *kservev1beta1.InferenceService)
+ }{
+ {
+ name: "predictor with user-defined node selector only",
+ modelSvc: &models.Service{
+ Name: "model-1", ModelName: "model", ModelVersion: "1", Namespace: project.Name,
+ ArtifactURI: "gs://my-artifacet", Type: models.ModelTypeTensorflow, Options: &models.ModelOption{},
+ Metadata: baseMeta, Protocol: protocol.HttpJson,
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ CPURequest: resource.MustParse("500m"), MemoryRequest: resource.MustParse("500Mi"),
+ NodeSelector: userNodeSelector,
+ },
+ },
+ deployConfig: baseDeployConfig,
+ checkFn: func(t *testing.T, infSvc *kservev1beta1.InferenceService) {
+ assert.Equal(t, userNodeSelector, infSvc.Spec.Predictor.NodeSelector)
+ },
+ },
+ {
+ name: "predictor with GPU node selector merged with user node selector",
+ modelSvc: &models.Service{
+ Name: "model-1", ModelName: "model", ModelVersion: "1", Namespace: project.Name,
+ ArtifactURI: "gs://my-artifacet", Type: models.ModelTypeTensorflow, Options: &models.ModelOption{},
+ Metadata: baseMeta, Protocol: protocol.HttpJson,
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ CPURequest: resource.MustParse("500m"), MemoryRequest: resource.MustParse("500Mi"),
+ GPUName: "NVIDIA T4", GPURequest: resource.MustParse("1"),
+ NodeSelector: userNodeSelector,
+ },
+ },
+ deployConfig: func() *config.DeploymentConfig {
+ cfg := *baseDeployConfig
+ cfg.GPUs = []config.GPUConfig{gpuConfig}
+ return &cfg
+ }(),
+ checkFn: func(t *testing.T, infSvc *kservev1beta1.InferenceService) {
+ assert.Equal(t, map[string]string{
+ "cloud.google.com/gke-accelerator": "nvidia-tesla-t4",
+ "pool": "workload-optimized",
+ }, infSvc.Spec.Predictor.NodeSelector)
+ // the shared GPU config must not be mutated by the merge
+ assert.Equal(t, gpuNodeSelector, gpuConfig.NodeSelector)
+ },
+ },
+ {
+ name: "transformer with user-defined node selector",
+ modelSvc: &models.Service{
+ Name: "model-1", ModelName: "model", ModelVersion: "1", Namespace: project.Name,
+ ArtifactURI: "gs://my-artifacet", Type: models.ModelTypeTensorflow, Options: &models.ModelOption{},
+ Metadata: baseMeta, Protocol: protocol.HttpJson,
+ Transformer: &models.Transformer{
+ Enabled: true, Image: "ghcr.io/gojek/merlin-transformer-test",
+ ResourceRequest: &models.ResourceRequest{
+ MinReplica: 1, MaxReplica: 2,
+ CPURequest: resource.MustParse("100m"), MemoryRequest: resource.MustParse("500Mi"),
+ NodeSelector: userNodeSelector,
+ },
+ },
+ },
+ deployConfig: baseDeployConfig,
+ checkFn: func(t *testing.T, infSvc *kservev1beta1.InferenceService) {
+ assert.NotNil(t, infSvc.Spec.Transformer)
+ assert.Equal(t, userNodeSelector, infSvc.Spec.Transformer.PodSpec.NodeSelector)
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ tpl := NewInferenceServiceTemplater(*tt.deployConfig)
+ infSvc, err := tpl.CreateInferenceServiceSpec(tt.modelSvc, defaultDeploymentScale)
+ assert.NoError(t, err)
+ tt.checkFn(t, infSvc)
+ })
+ }
+}
diff --git a/api/cmd/api/main.go b/api/cmd/api/main.go
index 6c5b54029..9be2acf31 100644
--- a/api/cmd/api/main.go
+++ b/api/cmd/api/main.go
@@ -286,6 +286,7 @@ func buildDependencies(ctx context.Context, cfg *config.Config, db *gorm.DB, dis
batchDeployment := initBatchDeployment(cfg, db, batchControllers, predJobBuilder)
predictionJobService := initPredictionJobService(cfg, batchControllers, predJobBuilder, db, dispatcher)
logService := initLogService(cfg)
+ nodePoolService := service.NewNodePoolService(clusterControllers)
// use "mlp" as product name for enforcer so that same policy can be reused by other components
enforcerCfg := enforcer.NewEnforcerBuilder().KetoEndpoints(cfg.AuthorizationConfig.KetoRemoteRead,
cfg.AuthorizationConfig.KetoRemoteWrite)
@@ -352,6 +353,7 @@ func buildDependencies(ctx context.Context, cfg *config.Config, db *gorm.DB, dis
VersionImageService: versionImageService,
EndpointsService: versionEndpointService,
LogService: logService,
+ NodePoolService: nodePoolService,
PredictionJobService: predictionJobService,
SecretService: secretService,
ModelEndpointAlertService: modelEndpointAlertService,
diff --git a/api/models/node_pool.go b/api/models/node_pool.go
new file mode 100644
index 000000000..3f2a6d02e
--- /dev/null
+++ b/api/models/node_pool.go
@@ -0,0 +1,106 @@
+// Copyright 2020 The Merlin Authors
+//
+// 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 models
+
+import (
+ "strings"
+
+ corev1 "k8s.io/api/core/v1"
+)
+
+// NodePool is a schedulable node pool discovered from the cluster. It pairs the
+// taint the pool's nodes carry (which a pod must tolerate) with the labels
+// common to those nodes (which a pod can use to pin onto the pool).
+type NodePool struct {
+ // Taint every node in the pool carries — drives the required toleration.
+ Taint corev1.Taint `json:"taint"`
+ // NodeSelector are labels shared by all nodes carrying the taint — drives pinning.
+ NodeSelector map[string]string `json:"node_selector"`
+}
+
+// volatileLabelPrefixes are Kubernetes/cloud-managed label prefixes that don't
+// identify a pool (hostname, zone, instance-type, etc.) and so are dropped from
+// the suggested node selector.
+var volatileLabelPrefixes = []string{
+ "kubernetes.io/",
+ "k8s.io/",
+ "node.kubernetes.io/",
+ "beta.kubernetes.io/",
+ "topology.kubernetes.io/",
+ "failure-domain.beta.kubernetes.io/",
+}
+
+// transientTaintPrefixes are taints Kubernetes adds/removes automatically; they
+// don't represent a deliberate pool and are excluded from discovery.
+var transientTaintPrefixes = []string{
+ "node.kubernetes.io/",
+ "node.cloudprovider.kubernetes.io/",
+}
+
+func hasAnyPrefix(s string, prefixes []string) bool {
+ for _, p := range prefixes {
+ if strings.HasPrefix(s, p) {
+ return true
+ }
+ }
+ return false
+}
+
+// AggregateNodePools groups nodes by the deliberate taints they carry and, for
+// each taint, computes the labels common to every node bearing it. The result
+// is the set of pools a model can target: tolerate the taint, pin with the
+// shared labels.
+func AggregateNodePools(nodes []corev1.Node) []NodePool {
+ type group struct {
+ taint corev1.Taint
+ labels map[string]string
+ count int
+ }
+
+ groups := map[string]*group{}
+ for _, node := range nodes {
+ for _, taint := range node.Spec.Taints {
+ if hasAnyPrefix(taint.Key, transientTaintPrefixes) {
+ continue
+ }
+ id := taint.Key + "=" + taint.Value + ":" + string(taint.Effect)
+ g, ok := groups[id]
+ if !ok {
+ // first node for this taint: seed the common-label set
+ labels := map[string]string{}
+ for k, v := range node.Labels {
+ if !hasAnyPrefix(k, volatileLabelPrefixes) {
+ labels[k] = v
+ }
+ }
+ groups[id] = &group{taint: taint, labels: labels, count: 1}
+ continue
+ }
+ // intersect: keep only labels present with the same value on this node too
+ for k, v := range g.labels {
+ if nv, ok := node.Labels[k]; !ok || nv != v {
+ delete(g.labels, k)
+ }
+ }
+ g.count++
+ }
+ }
+
+ pools := make([]NodePool, 0, len(groups))
+ for _, g := range groups {
+ pools = append(pools, NodePool{Taint: g.taint, NodeSelector: g.labels})
+ }
+ return pools
+}
diff --git a/api/models/node_pool_test.go b/api/models/node_pool_test.go
new file mode 100644
index 000000000..d4c1dd5b8
--- /dev/null
+++ b/api/models/node_pool_test.go
@@ -0,0 +1,68 @@
+package models
+
+import (
+ "sort"
+ "testing"
+
+ "github.com/stretchr/testify/assert"
+ corev1 "k8s.io/api/core/v1"
+ metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
+)
+
+func node(labels map[string]string, taints ...corev1.Taint) corev1.Node {
+ return corev1.Node{
+ ObjectMeta: metav1.ObjectMeta{Labels: labels},
+ Spec: corev1.NodeSpec{Taints: taints},
+ }
+}
+
+func TestAggregateNodePools(t *testing.T) {
+ woTaint := corev1.Taint{Key: "workload-optimized", Value: "true", Effect: corev1.TaintEffectNoSchedule}
+ transient := corev1.Taint{Key: "node.kubernetes.io/unreachable", Effect: corev1.TaintEffectNoExecute}
+
+ nodes := []corev1.Node{
+ // two nodes in the workload-optimized pool sharing pool + team labels,
+ // differing on hostname (volatile) — hostname must be dropped, team kept only if shared
+ node(map[string]string{
+ "pool": "workload-optimized",
+ "team": "dsp",
+ "kubernetes.io/hostname": "node-a",
+ }, woTaint, transient),
+ node(map[string]string{
+ "pool": "workload-optimized",
+ "team": "ml",
+ "kubernetes.io/hostname": "node-b",
+ }, woTaint),
+ // an untainted node — contributes no pool
+ node(map[string]string{"pool": "default"}),
+ }
+
+ pools := AggregateNodePools(nodes)
+
+ assert.Len(t, pools, 1, "only the deliberate workload-optimized taint should surface")
+ p := pools[0]
+ assert.Equal(t, woTaint, p.Taint)
+ // shared label kept; differing label (team) and volatile label (hostname) dropped
+ assert.Equal(t, map[string]string{"pool": "workload-optimized"}, p.NodeSelector)
+}
+
+func TestAggregateNodePools_MultiplePools(t *testing.T) {
+ a := corev1.Taint{Key: "pool-a", Value: "true", Effect: corev1.TaintEffectNoSchedule}
+ b := corev1.Taint{Key: "pool-b", Value: "true", Effect: corev1.TaintEffectNoSchedule}
+ nodes := []corev1.Node{
+ node(map[string]string{"pool": "a"}, a),
+ node(map[string]string{"pool": "b"}, b),
+ }
+
+ pools := AggregateNodePools(nodes)
+ keys := []string{}
+ for _, p := range pools {
+ keys = append(keys, p.Taint.Key)
+ }
+ sort.Strings(keys)
+ assert.Equal(t, []string{"pool-a", "pool-b"}, keys)
+}
+
+func TestAggregateNodePools_NoTaints(t *testing.T) {
+ assert.Empty(t, AggregateNodePools([]corev1.Node{node(map[string]string{"pool": "x"})}))
+}
diff --git a/api/models/resource_request.go b/api/models/resource_request.go
index 839641908..d370908da 100644
--- a/api/models/resource_request.go
+++ b/api/models/resource_request.go
@@ -44,6 +44,8 @@ type ResourceRequest struct {
ReadinessProbe *ProbeConfig `json:"readiness_probe,omitempty"`
// Tolerations allow the model pods to be scheduled onto nodes with matching taints
Tolerations []corev1.Toleration `json:"tolerations,omitempty"`
+ // NodeSelector pins the model pods onto nodes whose labels match every entry
+ NodeSelector map[string]string `json:"node_selector,omitempty"`
}
// ProbeConfig represents the configuration for Kubernetes liveness/readiness probes
diff --git a/api/service/node_pool_service.go b/api/service/node_pool_service.go
new file mode 100644
index 000000000..3d1ed7bdf
--- /dev/null
+++ b/api/service/node_pool_service.go
@@ -0,0 +1,53 @@
+// Copyright 2020 The Merlin Authors
+//
+// 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 service
+
+import (
+ "context"
+ "fmt"
+
+ "github.com/caraml-dev/merlin/cluster"
+ "github.com/caraml-dev/merlin/models"
+)
+
+// NodePoolService discovers the schedulable node pools of a cluster so users
+// can target dedicated / tainted nodes without knowing raw labels and taints.
+type NodePoolService interface {
+ ListNodePools(ctx context.Context, cluster string) ([]models.NodePool, error)
+}
+
+type nodePoolService struct {
+ clusterControllers map[string]cluster.Controller
+}
+
+// NewNodePoolService creates a NodePoolService.
+// clusterControllers is a map of cluster name to its cluster.Controller.
+func NewNodePoolService(clusterControllers map[string]cluster.Controller) NodePoolService {
+ return &nodePoolService{clusterControllers: clusterControllers}
+}
+
+func (s *nodePoolService) ListNodePools(ctx context.Context, clusterName string) ([]models.NodePool, error) {
+ controller, ok := s.clusterControllers[clusterName]
+ if !ok {
+ return nil, fmt.Errorf("unable to find cluster controller for cluster %s", clusterName)
+ }
+
+ nodeList, err := controller.ListNodes(ctx)
+ if err != nil {
+ return nil, fmt.Errorf("unable to list nodes in cluster %s: %w", clusterName, err)
+ }
+
+ return models.AggregateNodePools(nodeList.Items), nil
+}
diff --git a/swagger.yaml b/swagger.yaml
index 182a99e19..c184b71a7 100644
--- a/swagger.yaml
+++ b/swagger.yaml
@@ -67,6 +67,27 @@ paths:
type: array
items:
"$ref": "#/components/schemas/Environment"
+ "/environments/{environment_name}/node-pools":
+ get:
+ tags:
+ - environment
+ summary: List schedulable node pools for an environment
+ description: Returns the deliberate taints and their shared node labels, so a model can be pinned to (and tolerate) a dedicated node pool.
+ parameters:
+ - name: environment_name
+ in: path
+ required: true
+ schema:
+ type: string
+ responses:
+ "200":
+ description: OK
+ content:
+ "*/*":
+ schema:
+ type: array
+ items:
+ "$ref": "#/components/schemas/NodePool"
"/projects":
get:
tags:
@@ -2126,6 +2147,30 @@ components:
description: Tolerations allow model pods to be scheduled onto nodes with matching taints
items:
"$ref": "#/components/schemas/Toleration"
+ node_selector:
+ type: object
+ description: NodeSelector pins model pods onto nodes whose labels match every entry
+ additionalProperties:
+ type: string
+ NodePool:
+ type: object
+ description: A schedulable node pool discovered from the cluster's node taints and shared labels
+ properties:
+ taint:
+ type: object
+ description: The taint the pool's nodes carry; a pod must tolerate it to run there
+ properties:
+ key:
+ type: string
+ value:
+ type: string
+ effect:
+ type: string
+ node_selector:
+ type: object
+ description: Labels shared by all nodes carrying the taint; use as node selector to pin onto the pool
+ additionalProperties:
+ type: string
Toleration:
type: object
description: Kubernetes toleration for scheduling pods onto tainted nodes
diff --git a/ui/src/components/ResourcesConfigTable.js b/ui/src/components/ResourcesConfigTable.js
index 49661a9e0..6dcee7db5 100644
--- a/ui/src/components/ResourcesConfigTable.js
+++ b/ui/src/components/ResourcesConfigTable.js
@@ -30,6 +30,7 @@ export const ResourcesConfigTable = ({
liveness_probe,
readiness_probe,
tolerations,
+ node_selector,
},
}) => {
const items = [
@@ -120,6 +121,16 @@ export const ResourcesConfigTable = ({
});
}
+ // Add node selectors if configured
+ if (node_selector && Object.keys(node_selector).length > 0) {
+ Object.entries(node_selector).forEach(([key, value], idx) => {
+ items.push({
+ title: idx === 0 ? "Node Selector" : "",
+ description: `${key}: ${value}`,
+ });
+ });
+ }
+
return (
+ `${pool.taint.key}=${pool.taint.value || ""}:${pool.taint.effect || ""}`;
+
+/**
+ * NodePoolSelect - a single dropdown that pins a model to a discovered node pool.
+ *
+ * Selecting a pool sets the node selector (to pin) and adds the pool's matching
+ * toleration (to permit), so the user never types raw labels or taints.
+ *
+ * Props
+ * environment – environment name (used to fetch pools for its cluster)
+ * nodeSelector – current node_selector object
+ * tolerations – current tolerations array
+ * onSelect(sel, tols) – called with the new node_selector + tolerations
+ */
+export const NodePoolSelect = ({
+ environment,
+ nodeSelector = {},
+ tolerations = [],
+ onSelect,
+}) => {
+ const [selected, setSelected] = useState(DEFAULT_VALUE);
+
+ const [{ data: pools, isLoaded }] = useMerlinApi(
+ `/environments/${environment}/node-pools`,
+ {},
+ [],
+ !!environment
+ );
+
+ const options = [
+ {
+ value: DEFAULT_VALUE,
+ inputDisplay: "Default — any available node",
+ },
+ ...(pools || []).map((pool) => ({
+ value: poolId(pool),
+ inputDisplay: `${pool.taint.key}=${pool.taint.value || "\"\""}`,
+ dropdownDisplay: (
+
+ {`${pool.taint.key}=${pool.taint.value || "\"\""}`}
+
+ {`taint ${poolId(pool)} · selector ${JSON.stringify(pool.node_selector || {})}`}
+
+
+ ),
+ })),
+ ];
+
+ const onChange = (value) => {
+ setSelected(value);
+ if (value === DEFAULT_VALUE) {
+ onSelect({}, []);
+ return;
+ }
+ const pool = (pools || []).find((p) => poolId(p) === value);
+ if (!pool) return;
+ const toleration = {
+ key: pool.taint.key,
+ operator: "Equal",
+ ...(pool.taint.value ? { value: pool.taint.value } : {}),
+ ...(pool.taint.effect ? { effect: pool.taint.effect } : {}),
+ };
+ onSelect({ ...(pool.node_selector || {}) }, [toleration]);
+ };
+
+ return (
+
+ );
+};
diff --git a/ui/src/pages/version/components/forms/components/NodeSelectorFormGroup.js b/ui/src/pages/version/components/forms/components/NodeSelectorFormGroup.js
new file mode 100644
index 000000000..2985d2fd4
--- /dev/null
+++ b/ui/src/pages/version/components/forms/components/NodeSelectorFormGroup.js
@@ -0,0 +1,100 @@
+import React, { Fragment, useState } from "react";
+import {
+ EuiButtonEmpty,
+ EuiButtonIcon,
+ EuiDescribedFormGroup,
+ EuiFieldText,
+ EuiFlexGroup,
+ EuiFlexItem,
+ EuiFormRow,
+ EuiSpacer,
+ EuiText,
+} from "@elastic/eui";
+
+const toRows = (obj) => Object.entries(obj || {}).map(([key, value]) => ({ key, value }));
+
+const toObject = (rows) => {
+ const obj = {};
+ rows.forEach((r) => {
+ const key = (r.key || "").trim();
+ if (key !== "") obj[key] = r.value || "";
+ });
+ return obj;
+};
+
+/**
+ * NodeSelectorFormGroup - edit the pod nodeSelector as a list of label key/value pairs.
+ *
+ * Props
+ * nodeSelector – object of label key -> value (may be undefined / {})
+ * onChangeHandler – called with the rebuilt object whenever the list changes
+ */
+export const NodeSelectorFormGroup = ({ nodeSelector = {}, onChangeHandler }) => {
+ const [rows, setRows] = useState(() => toRows(nodeSelector));
+
+ const push = (next) => {
+ setRows(next);
+ onChangeHandler(toObject(next));
+ };
+
+ const onChangeCell = (idx, field) => (e) => {
+ const value = e.target.value;
+ push(rows.map((r, i) => (i === idx ? { ...r, [field]: value } : r)));
+ };
+
+ const onAddRow = () => push([...rows, { key: "", value: "" }]);
+
+ const onDeleteRow = (idx) => () => push(rows.filter((_, i) => i !== idx));
+
+ return (
+ Node Selector}
+ description={
+
+
+ Pin the pods onto nodes whose labels match every entry. Combine with
+ a matching toleration to run on dedicated / tainted node pools.
+
+
+ }
+ fullWidth
+ >
+
+ {rows.map((row, idx) => (
+
+
+
+
+
+
+
+
+
+
+
+ ))}
+
+
+
+ Add node selector
+
+
+
+ );
+};
diff --git a/ui/src/pages/version/components/forms/steps/ModelStep.js b/ui/src/pages/version/components/forms/steps/ModelStep.js
index 7bf977cf8..02b902f04 100644
--- a/ui/src/pages/version/components/forms/steps/ModelStep.js
+++ b/ui/src/pages/version/components/forms/steps/ModelStep.js
@@ -4,7 +4,7 @@ import {
get,
useOnChangeHandler,
} from "@caraml-dev/ui-lib";
-import { EuiAccordion, EuiFlexGroup, EuiFlexItem, EuiSpacer } from "@elastic/eui";
+import { EuiAccordion, EuiDescribedFormGroup, EuiFlexGroup, EuiFlexItem, EuiSpacer } from "@elastic/eui";
import React, { useContext } from "react";
import { PROTOCOL } from "../../../../../services/version_endpoint/VersionEndpoint";
import { DeploymentConfigPanel } from "../components/DeploymentConfigPanel";
@@ -16,6 +16,8 @@ import { ImageBuilderSection } from "../components/ImageBuilderSection";
import { CPULimitsFormGroup } from "../components/CPULimitsFormGroup";
import { ProbesFormGroup } from "../components/ProbesFormGroup";
import { TolerationFormGroup } from "../components/TolerationFormGroup";
+import { NodeSelectorFormGroup } from "../components/NodeSelectorFormGroup";
+import { NodePoolSelect } from "../components/NodePoolSelect";
export const ModelStep = ({ version, isEnvironmentDisabled = false, maxAllowedReplica, setMaxAllowedReplica }) => {
const { data, onChangeHandler } = useContext(FormContext);
@@ -67,11 +69,32 @@ export const ModelStep = ({ version, isEnvironmentDisabled = false, maxAllowedRe
errors={get(errors, "image_builder_resource_request")}
/>
+ Node Pool}
+ description="Pick a pool to pin this model onto and tolerate — sets the node selector and toleration for you."
+ fullWidth
+ >
+ {
+ onChange("resource_request.node_selector")(sel);
+ onChange("resource_request.tolerations")(tols);
+ }}
+ />
+
+
+
+
}
/>
diff --git a/ui/src/pages/version/components/forms/steps/TransformerStep.js b/ui/src/pages/version/components/forms/steps/TransformerStep.js
index 40305ba02..07c7ca3a6 100644
--- a/ui/src/pages/version/components/forms/steps/TransformerStep.js
+++ b/ui/src/pages/version/components/forms/steps/TransformerStep.js
@@ -14,6 +14,7 @@ import { ResourcesPanel } from "../components/ResourcesPanel";
import { SelectTransformerPanel } from "../components/SelectTransformerPanel";
import { CPULimitsFormGroup } from "../components/CPULimitsFormGroup";
import { TolerationFormGroup } from "../components/TolerationFormGroup";
+import { NodeSelectorFormGroup } from "../components/NodeSelectorFormGroup";
export const TransformerStep = ({ maxAllowedReplica }) => {
const {
@@ -59,6 +60,11 @@ export const TransformerStep = ({ maxAllowedReplica }) => {
onChangeHandler={onChange("transformer.resource_request.tolerations")}
errors={get(errors, "transformer.resource_request.tolerations")}
/>
+
+
}
/>
diff --git a/ui/src/services/transformer/Transformer.js b/ui/src/services/transformer/Transformer.js
index b827f4bd0..c90745b30 100644
--- a/ui/src/services/transformer/Transformer.js
+++ b/ui/src/services/transformer/Transformer.js
@@ -25,6 +25,7 @@ export class Transformer {
cpu_limit: "",
memory_request: "512Mi",
tolerations: [],
+ node_selector: {},
};
this.env_vars = [];
@@ -65,6 +66,10 @@ export class Transformer {
transformer.resource_request.tolerations = [];
}
+ if (transformer.resource_request && !transformer.resource_request.node_selector) {
+ transformer.resource_request.node_selector = {};
+ }
+
return transformer;
}
diff --git a/ui/src/services/version_endpoint/VersionEndpoint.js b/ui/src/services/version_endpoint/VersionEndpoint.js
index 40c792ec1..3b45a5e5a 100644
--- a/ui/src/services/version_endpoint/VersionEndpoint.js
+++ b/ui/src/services/version_endpoint/VersionEndpoint.js
@@ -31,6 +31,7 @@ export class VersionEndpoint {
liveness_probe: null,
readiness_probe: null,
tolerations: [],
+ node_selector: {},
};
this.image_builder_resource_request = {
@@ -79,6 +80,10 @@ export class VersionEndpoint {
versionEndpoint.resource_request.tolerations = [];
}
+ if (versionEndpoint.resource_request && !versionEndpoint.resource_request.node_selector) {
+ versionEndpoint.resource_request.node_selector = {};
+ }
+
if (json.transformer) {
versionEndpoint.transformer = Transformer.fromJson(json.transformer);
}
From e5ce4cccade32d9f4fa21322b0ad525bfb755eea Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Mon, 20 Jul 2026 12:44:52 +0530
Subject: [PATCH 12/14] added cross cluster node fetch.
---
api/api/node_pool_api.go | 5 ++++-
api/service/node_pool_service.go | 20 ++++++++++++--------
2 files changed, 16 insertions(+), 9 deletions(-)
diff --git a/api/api/node_pool_api.go b/api/api/node_pool_api.go
index d355d18b2..76b258623 100644
--- a/api/api/node_pool_api.go
+++ b/api/api/node_pool_api.go
@@ -41,7 +41,10 @@ func (c *NodePoolController) ListNodePools(r *http.Request, vars map[string]stri
return InternalServerError(fmt.Sprintf("Error getting environment: %v", err))
}
- nodePools, err := c.NodePoolService.ListNodePools(ctx, env.Cluster)
+ // Route by environment name — this selects the same per-environment controller
+ // that deployment uses, so nodes are read from the exact cluster the model
+ // deploys to (in-cluster SA locally, or mTLS client cert for a remote cluster).
+ nodePools, err := c.NodePoolService.ListNodePools(ctx, env.Name)
if err != nil {
return InternalServerError(fmt.Sprintf("Error listing node pools: %v", err))
}
diff --git a/api/service/node_pool_service.go b/api/service/node_pool_service.go
index 3d1ed7bdf..8d28dbf80 100644
--- a/api/service/node_pool_service.go
+++ b/api/service/node_pool_service.go
@@ -22,31 +22,35 @@ import (
"github.com/caraml-dev/merlin/models"
)
-// NodePoolService discovers the schedulable node pools of a cluster so users
-// can target dedicated / tainted nodes without knowing raw labels and taints.
+// NodePoolService discovers the schedulable node pools of the cluster backing a
+// deployment environment, so users can target dedicated / tainted nodes without
+// knowing raw labels and taints. It reuses the same per-environment controller
+// that deployment uses, so nodes are always read from the exact cluster (and via
+// the exact credentials, e.g. mTLS for remote clusters) the model deploys to.
type NodePoolService interface {
- ListNodePools(ctx context.Context, cluster string) ([]models.NodePool, error)
+ ListNodePools(ctx context.Context, environmentName string) ([]models.NodePool, error)
}
type nodePoolService struct {
+ // clusterControllers is a map of environment name to its cluster.Controller.
clusterControllers map[string]cluster.Controller
}
// NewNodePoolService creates a NodePoolService.
-// clusterControllers is a map of cluster name to its cluster.Controller.
+// clusterControllers is a map of environment name to its cluster.Controller.
func NewNodePoolService(clusterControllers map[string]cluster.Controller) NodePoolService {
return &nodePoolService{clusterControllers: clusterControllers}
}
-func (s *nodePoolService) ListNodePools(ctx context.Context, clusterName string) ([]models.NodePool, error) {
- controller, ok := s.clusterControllers[clusterName]
+func (s *nodePoolService) ListNodePools(ctx context.Context, environmentName string) ([]models.NodePool, error) {
+ controller, ok := s.clusterControllers[environmentName]
if !ok {
- return nil, fmt.Errorf("unable to find cluster controller for cluster %s", clusterName)
+ return nil, fmt.Errorf("unable to find cluster controller for environment %s", environmentName)
}
nodeList, err := controller.ListNodes(ctx)
if err != nil {
- return nil, fmt.Errorf("unable to list nodes in cluster %s: %w", clusterName, err)
+ return nil, fmt.Errorf("unable to list nodes for environment %s: %w", environmentName, err)
}
return models.AggregateNodePools(nodeList.Items), nil
From 4ce562fb34b2e72d461fc16dbd71c918e4fde0ef Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Mon, 20 Jul 2026 14:28:45 +0530
Subject: [PATCH 13/14] added node list and selector
---
api/api/{node_pool_api.go => node_api.go} | 16 +--
api/api/router.go | 6 +-
api/cmd/api/main.go | 4 +-
api/models/node.go | 51 +++++++++
api/models/node_pool.go | 106 ------------------
api/models/node_pool_test.go | 68 -----------
api/models/node_test.go | 48 ++++++++
.../{node_pool_service.go => node_service.go} | 26 ++---
swagger.yaml | 42 ++++---
.../forms/components/NodePoolSelect.js | 83 --------------
.../components/forms/steps/ModelStep.js | 13 ++-
.../components/forms/steps/TransformerStep.js | 16 +--
12 files changed, 159 insertions(+), 320 deletions(-)
rename api/api/{node_pool_api.go => node_api.go} (70%)
create mode 100644 api/models/node.go
delete mode 100644 api/models/node_pool.go
delete mode 100644 api/models/node_pool_test.go
create mode 100644 api/models/node_test.go
rename api/service/{node_pool_service.go => node_service.go} (57%)
delete mode 100644 ui/src/pages/version/components/forms/components/NodePoolSelect.js
diff --git a/api/api/node_pool_api.go b/api/api/node_api.go
similarity index 70%
rename from api/api/node_pool_api.go
rename to api/api/node_api.go
index 76b258623..253ebe3a7 100644
--- a/api/api/node_pool_api.go
+++ b/api/api/node_api.go
@@ -22,14 +22,14 @@ import (
"gorm.io/gorm"
)
-// NodePoolController serves node pool discovery for a deployment environment.
-type NodePoolController struct {
+// NodeController serves node listing for a deployment environment.
+type NodeController struct {
*AppContext
}
-// ListNodePools returns the schedulable node pools (taint + shared labels) of the
-// cluster backing the given environment, so the UI can offer them for placement.
-func (c *NodePoolController) ListNodePools(r *http.Request, vars map[string]string, _ interface{}) *Response {
+// ListNodes returns the nodes (name, status, labels, taints) of the cluster
+// backing the given environment, so the UI can offer them for model placement.
+func (c *NodeController) ListNodes(r *http.Request, vars map[string]string, _ interface{}) *Response {
ctx := r.Context()
environmentName := vars["environment_name"]
@@ -44,10 +44,10 @@ func (c *NodePoolController) ListNodePools(r *http.Request, vars map[string]stri
// Route by environment name — this selects the same per-environment controller
// that deployment uses, so nodes are read from the exact cluster the model
// deploys to (in-cluster SA locally, or mTLS client cert for a remote cluster).
- nodePools, err := c.NodePoolService.ListNodePools(ctx, env.Name)
+ nodes, err := c.NodeService.ListNodes(ctx, env.Name)
if err != nil {
- return InternalServerError(fmt.Sprintf("Error listing node pools: %v", err))
+ return InternalServerError(fmt.Sprintf("Error listing nodes: %v", err))
}
- return Ok(nodePools)
+ return Ok(nodes)
}
diff --git a/api/api/router.go b/api/api/router.go
index 5b584bf22..c548f734e 100644
--- a/api/api/router.go
+++ b/api/api/router.go
@@ -62,7 +62,7 @@ type AppContext struct {
VersionImageService service.VersionImageService
EndpointsService service.EndpointsService
LogService service.LogService
- NodePoolService service.NodePoolService
+ NodeService service.NodeService
PredictionJobService service.PredictionJobService
SecretService service.SecretService
ModelEndpointAlertService service.ModelEndpointAlertService
@@ -169,7 +169,7 @@ func NewRouter(appCtx AppContext) (*mux.Router, error) {
endpointsController := EndpointsController{&appCtx}
predictionJobController := PredictionJobController{&appCtx}
logController := LogController{&appCtx}
- nodePoolController := NodePoolController{&appCtx}
+ nodeController := NodeController{&appCtx}
secretController := SecretsController{&appCtx}
alertsController := AlertsController{&appCtx}
transformerController := TransformerController{&appCtx}
@@ -178,7 +178,7 @@ func NewRouter(appCtx AppContext) (*mux.Router, error) {
routes := []Route{
// Environment API
{http.MethodGet, "/environments", nil, environmentController.ListEnvironments, "ListEnvironments"},
- {http.MethodGet, "/environments/{environment_name}/node-pools", nil, nodePoolController.ListNodePools, "ListNodePools"},
+ {http.MethodGet, "/environments/{environment_name}/nodes", nil, nodeController.ListNodes, "ListNodes"},
// Project API
{http.MethodGet, "/projects/{project_id:[0-9]+}", nil, projectsController.GetProject, "GetProject"},
diff --git a/api/cmd/api/main.go b/api/cmd/api/main.go
index 9be2acf31..072b34956 100644
--- a/api/cmd/api/main.go
+++ b/api/cmd/api/main.go
@@ -286,7 +286,7 @@ func buildDependencies(ctx context.Context, cfg *config.Config, db *gorm.DB, dis
batchDeployment := initBatchDeployment(cfg, db, batchControllers, predJobBuilder)
predictionJobService := initPredictionJobService(cfg, batchControllers, predJobBuilder, db, dispatcher)
logService := initLogService(cfg)
- nodePoolService := service.NewNodePoolService(clusterControllers)
+ nodeService := service.NewNodeService(clusterControllers)
// use "mlp" as product name for enforcer so that same policy can be reused by other components
enforcerCfg := enforcer.NewEnforcerBuilder().KetoEndpoints(cfg.AuthorizationConfig.KetoRemoteRead,
cfg.AuthorizationConfig.KetoRemoteWrite)
@@ -353,7 +353,7 @@ func buildDependencies(ctx context.Context, cfg *config.Config, db *gorm.DB, dis
VersionImageService: versionImageService,
EndpointsService: versionEndpointService,
LogService: logService,
- NodePoolService: nodePoolService,
+ NodeService: nodeService,
PredictionJobService: predictionJobService,
SecretService: secretService,
ModelEndpointAlertService: modelEndpointAlertService,
diff --git a/api/models/node.go b/api/models/node.go
new file mode 100644
index 000000000..4391cc559
--- /dev/null
+++ b/api/models/node.go
@@ -0,0 +1,51 @@
+// Copyright 2020 The Merlin Authors
+//
+// 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 models
+
+import (
+ corev1 "k8s.io/api/core/v1"
+)
+
+// NodeInfo is a summary of a cluster node relevant to model placement: its
+// labels (usable as a node selector) and taints (which a pod must tolerate).
+type NodeInfo struct {
+ Name string `json:"name"`
+ Ready bool `json:"ready"`
+ Labels map[string]string `json:"labels"`
+ Taints []corev1.Taint `json:"taints"`
+}
+
+func isNodeReady(node corev1.Node) bool {
+ for _, cond := range node.Status.Conditions {
+ if cond.Type == corev1.NodeReady {
+ return cond.Status == corev1.ConditionTrue
+ }
+ }
+ return false
+}
+
+// NewNodeInfos maps Kubernetes nodes to placement-relevant summaries.
+func NewNodeInfos(nodes []corev1.Node) []NodeInfo {
+ infos := make([]NodeInfo, 0, len(nodes))
+ for _, node := range nodes {
+ infos = append(infos, NodeInfo{
+ Name: node.Name,
+ Ready: isNodeReady(node),
+ Labels: node.Labels,
+ Taints: node.Spec.Taints,
+ })
+ }
+ return infos
+}
diff --git a/api/models/node_pool.go b/api/models/node_pool.go
deleted file mode 100644
index 3f2a6d02e..000000000
--- a/api/models/node_pool.go
+++ /dev/null
@@ -1,106 +0,0 @@
-// Copyright 2020 The Merlin Authors
-//
-// 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 models
-
-import (
- "strings"
-
- corev1 "k8s.io/api/core/v1"
-)
-
-// NodePool is a schedulable node pool discovered from the cluster. It pairs the
-// taint the pool's nodes carry (which a pod must tolerate) with the labels
-// common to those nodes (which a pod can use to pin onto the pool).
-type NodePool struct {
- // Taint every node in the pool carries — drives the required toleration.
- Taint corev1.Taint `json:"taint"`
- // NodeSelector are labels shared by all nodes carrying the taint — drives pinning.
- NodeSelector map[string]string `json:"node_selector"`
-}
-
-// volatileLabelPrefixes are Kubernetes/cloud-managed label prefixes that don't
-// identify a pool (hostname, zone, instance-type, etc.) and so are dropped from
-// the suggested node selector.
-var volatileLabelPrefixes = []string{
- "kubernetes.io/",
- "k8s.io/",
- "node.kubernetes.io/",
- "beta.kubernetes.io/",
- "topology.kubernetes.io/",
- "failure-domain.beta.kubernetes.io/",
-}
-
-// transientTaintPrefixes are taints Kubernetes adds/removes automatically; they
-// don't represent a deliberate pool and are excluded from discovery.
-var transientTaintPrefixes = []string{
- "node.kubernetes.io/",
- "node.cloudprovider.kubernetes.io/",
-}
-
-func hasAnyPrefix(s string, prefixes []string) bool {
- for _, p := range prefixes {
- if strings.HasPrefix(s, p) {
- return true
- }
- }
- return false
-}
-
-// AggregateNodePools groups nodes by the deliberate taints they carry and, for
-// each taint, computes the labels common to every node bearing it. The result
-// is the set of pools a model can target: tolerate the taint, pin with the
-// shared labels.
-func AggregateNodePools(nodes []corev1.Node) []NodePool {
- type group struct {
- taint corev1.Taint
- labels map[string]string
- count int
- }
-
- groups := map[string]*group{}
- for _, node := range nodes {
- for _, taint := range node.Spec.Taints {
- if hasAnyPrefix(taint.Key, transientTaintPrefixes) {
- continue
- }
- id := taint.Key + "=" + taint.Value + ":" + string(taint.Effect)
- g, ok := groups[id]
- if !ok {
- // first node for this taint: seed the common-label set
- labels := map[string]string{}
- for k, v := range node.Labels {
- if !hasAnyPrefix(k, volatileLabelPrefixes) {
- labels[k] = v
- }
- }
- groups[id] = &group{taint: taint, labels: labels, count: 1}
- continue
- }
- // intersect: keep only labels present with the same value on this node too
- for k, v := range g.labels {
- if nv, ok := node.Labels[k]; !ok || nv != v {
- delete(g.labels, k)
- }
- }
- g.count++
- }
- }
-
- pools := make([]NodePool, 0, len(groups))
- for _, g := range groups {
- pools = append(pools, NodePool{Taint: g.taint, NodeSelector: g.labels})
- }
- return pools
-}
diff --git a/api/models/node_pool_test.go b/api/models/node_pool_test.go
deleted file mode 100644
index d4c1dd5b8..000000000
--- a/api/models/node_pool_test.go
+++ /dev/null
@@ -1,68 +0,0 @@
-package models
-
-import (
- "sort"
- "testing"
-
- "github.com/stretchr/testify/assert"
- corev1 "k8s.io/api/core/v1"
- metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
-)
-
-func node(labels map[string]string, taints ...corev1.Taint) corev1.Node {
- return corev1.Node{
- ObjectMeta: metav1.ObjectMeta{Labels: labels},
- Spec: corev1.NodeSpec{Taints: taints},
- }
-}
-
-func TestAggregateNodePools(t *testing.T) {
- woTaint := corev1.Taint{Key: "workload-optimized", Value: "true", Effect: corev1.TaintEffectNoSchedule}
- transient := corev1.Taint{Key: "node.kubernetes.io/unreachable", Effect: corev1.TaintEffectNoExecute}
-
- nodes := []corev1.Node{
- // two nodes in the workload-optimized pool sharing pool + team labels,
- // differing on hostname (volatile) — hostname must be dropped, team kept only if shared
- node(map[string]string{
- "pool": "workload-optimized",
- "team": "dsp",
- "kubernetes.io/hostname": "node-a",
- }, woTaint, transient),
- node(map[string]string{
- "pool": "workload-optimized",
- "team": "ml",
- "kubernetes.io/hostname": "node-b",
- }, woTaint),
- // an untainted node — contributes no pool
- node(map[string]string{"pool": "default"}),
- }
-
- pools := AggregateNodePools(nodes)
-
- assert.Len(t, pools, 1, "only the deliberate workload-optimized taint should surface")
- p := pools[0]
- assert.Equal(t, woTaint, p.Taint)
- // shared label kept; differing label (team) and volatile label (hostname) dropped
- assert.Equal(t, map[string]string{"pool": "workload-optimized"}, p.NodeSelector)
-}
-
-func TestAggregateNodePools_MultiplePools(t *testing.T) {
- a := corev1.Taint{Key: "pool-a", Value: "true", Effect: corev1.TaintEffectNoSchedule}
- b := corev1.Taint{Key: "pool-b", Value: "true", Effect: corev1.TaintEffectNoSchedule}
- nodes := []corev1.Node{
- node(map[string]string{"pool": "a"}, a),
- node(map[string]string{"pool": "b"}, b),
- }
-
- pools := AggregateNodePools(nodes)
- keys := []string{}
- for _, p := range pools {
- keys = append(keys, p.Taint.Key)
- }
- sort.Strings(keys)
- assert.Equal(t, []string{"pool-a", "pool-b"}, keys)
-}
-
-func TestAggregateNodePools_NoTaints(t *testing.T) {
- assert.Empty(t, AggregateNodePools([]corev1.Node{node(map[string]string{"pool": "x"})}))
-}
diff --git a/api/models/node_test.go b/api/models/node_test.go
new file mode 100644
index 000000000..21e963899
--- /dev/null
+++ b/api/models/node_test.go
@@ -0,0 +1,48 @@
+package models
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/assert"
+ corev1 "k8s.io/api/core/v1"
+ metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
+)
+
+func TestNewNodeInfos(t *testing.T) {
+ nodes := []corev1.Node{
+ {
+ ObjectMeta: metav1.ObjectMeta{
+ Name: "node-a",
+ Labels: map[string]string{"pool": "workload-optimized"},
+ },
+ Spec: corev1.NodeSpec{
+ Taints: []corev1.Taint{
+ {Key: "workload-optimized", Value: "true", Effect: corev1.TaintEffectNoSchedule},
+ },
+ },
+ Status: corev1.NodeStatus{
+ Conditions: []corev1.NodeCondition{
+ {Type: corev1.NodeReady, Status: corev1.ConditionTrue},
+ },
+ },
+ },
+ {
+ ObjectMeta: metav1.ObjectMeta{Name: "node-b"},
+ Status: corev1.NodeStatus{
+ Conditions: []corev1.NodeCondition{
+ {Type: corev1.NodeReady, Status: corev1.ConditionFalse},
+ },
+ },
+ },
+ }
+
+ infos := NewNodeInfos(nodes)
+
+ assert.Len(t, infos, 2)
+ assert.Equal(t, "node-a", infos[0].Name)
+ assert.True(t, infos[0].Ready)
+ assert.Equal(t, map[string]string{"pool": "workload-optimized"}, infos[0].Labels)
+ assert.Len(t, infos[0].Taints, 1)
+ assert.Equal(t, "node-b", infos[1].Name)
+ assert.False(t, infos[1].Ready)
+}
diff --git a/api/service/node_pool_service.go b/api/service/node_service.go
similarity index 57%
rename from api/service/node_pool_service.go
rename to api/service/node_service.go
index 8d28dbf80..42e645c97 100644
--- a/api/service/node_pool_service.go
+++ b/api/service/node_service.go
@@ -22,27 +22,27 @@ import (
"github.com/caraml-dev/merlin/models"
)
-// NodePoolService discovers the schedulable node pools of the cluster backing a
-// deployment environment, so users can target dedicated / tainted nodes without
-// knowing raw labels and taints. It reuses the same per-environment controller
-// that deployment uses, so nodes are always read from the exact cluster (and via
-// the exact credentials, e.g. mTLS for remote clusters) the model deploys to.
-type NodePoolService interface {
- ListNodePools(ctx context.Context, environmentName string) ([]models.NodePool, error)
+// NodeService lists the nodes of the cluster backing a deployment environment,
+// so users can see the available nodes (labels + taints) for model placement.
+// It reuses the same per-environment controller that deployment uses, so nodes
+// are read from the exact cluster (and via the exact credentials, e.g. mTLS for
+// remote clusters) the model deploys to.
+type NodeService interface {
+ ListNodes(ctx context.Context, environmentName string) ([]models.NodeInfo, error)
}
-type nodePoolService struct {
+type nodeService struct {
// clusterControllers is a map of environment name to its cluster.Controller.
clusterControllers map[string]cluster.Controller
}
-// NewNodePoolService creates a NodePoolService.
+// NewNodeService creates a NodeService.
// clusterControllers is a map of environment name to its cluster.Controller.
-func NewNodePoolService(clusterControllers map[string]cluster.Controller) NodePoolService {
- return &nodePoolService{clusterControllers: clusterControllers}
+func NewNodeService(clusterControllers map[string]cluster.Controller) NodeService {
+ return &nodeService{clusterControllers: clusterControllers}
}
-func (s *nodePoolService) ListNodePools(ctx context.Context, environmentName string) ([]models.NodePool, error) {
+func (s *nodeService) ListNodes(ctx context.Context, environmentName string) ([]models.NodeInfo, error) {
controller, ok := s.clusterControllers[environmentName]
if !ok {
return nil, fmt.Errorf("unable to find cluster controller for environment %s", environmentName)
@@ -53,5 +53,5 @@ func (s *nodePoolService) ListNodePools(ctx context.Context, environmentName str
return nil, fmt.Errorf("unable to list nodes for environment %s: %w", environmentName, err)
}
- return models.AggregateNodePools(nodeList.Items), nil
+ return models.NewNodeInfos(nodeList.Items), nil
}
diff --git a/swagger.yaml b/swagger.yaml
index c184b71a7..f66d40b89 100644
--- a/swagger.yaml
+++ b/swagger.yaml
@@ -67,12 +67,12 @@ paths:
type: array
items:
"$ref": "#/components/schemas/Environment"
- "/environments/{environment_name}/node-pools":
+ "/environments/{environment_name}/nodes":
get:
tags:
- environment
- summary: List schedulable node pools for an environment
- description: Returns the deliberate taints and their shared node labels, so a model can be pinned to (and tolerate) a dedicated node pool.
+ summary: List nodes for an environment
+ description: Returns the nodes (name, status, labels, taints) of the cluster backing the environment, for use in model placement (node selector + tolerations).
parameters:
- name: environment_name
in: path
@@ -87,7 +87,7 @@ paths:
schema:
type: array
items:
- "$ref": "#/components/schemas/NodePool"
+ "$ref": "#/components/schemas/NodeInfo"
"/projects":
get:
tags:
@@ -2152,25 +2152,31 @@ components:
description: NodeSelector pins model pods onto nodes whose labels match every entry
additionalProperties:
type: string
- NodePool:
+ NodeInfo:
type: object
- description: A schedulable node pool discovered from the cluster's node taints and shared labels
+ description: A cluster node's placement-relevant summary
properties:
- taint:
- type: object
- description: The taint the pool's nodes carry; a pod must tolerate it to run there
- properties:
- key:
- type: string
- value:
- type: string
- effect:
- type: string
- node_selector:
+ name:
+ type: string
+ ready:
+ type: boolean
+ labels:
type: object
- description: Labels shared by all nodes carrying the taint; use as node selector to pin onto the pool
+ description: Node labels; usable as a node selector to pin a pod onto this node's pool
additionalProperties:
type: string
+ taints:
+ type: array
+ description: Node taints; a pod must tolerate these to be scheduled here
+ items:
+ type: object
+ properties:
+ key:
+ type: string
+ value:
+ type: string
+ effect:
+ type: string
Toleration:
type: object
description: Kubernetes toleration for scheduling pods onto tainted nodes
diff --git a/ui/src/pages/version/components/forms/components/NodePoolSelect.js b/ui/src/pages/version/components/forms/components/NodePoolSelect.js
deleted file mode 100644
index 97c38cbc5..000000000
--- a/ui/src/pages/version/components/forms/components/NodePoolSelect.js
+++ /dev/null
@@ -1,83 +0,0 @@
-import React, { Fragment, useState } from "react";
-import { EuiSuperSelect, EuiText } from "@elastic/eui";
-import { useMerlinApi } from "../../../../../hooks/useMerlinApi";
-
-const DEFAULT_VALUE = "__default__";
-
-const poolId = (pool) =>
- `${pool.taint.key}=${pool.taint.value || ""}:${pool.taint.effect || ""}`;
-
-/**
- * NodePoolSelect - a single dropdown that pins a model to a discovered node pool.
- *
- * Selecting a pool sets the node selector (to pin) and adds the pool's matching
- * toleration (to permit), so the user never types raw labels or taints.
- *
- * Props
- * environment – environment name (used to fetch pools for its cluster)
- * nodeSelector – current node_selector object
- * tolerations – current tolerations array
- * onSelect(sel, tols) – called with the new node_selector + tolerations
- */
-export const NodePoolSelect = ({
- environment,
- nodeSelector = {},
- tolerations = [],
- onSelect,
-}) => {
- const [selected, setSelected] = useState(DEFAULT_VALUE);
-
- const [{ data: pools, isLoaded }] = useMerlinApi(
- `/environments/${environment}/node-pools`,
- {},
- [],
- !!environment
- );
-
- const options = [
- {
- value: DEFAULT_VALUE,
- inputDisplay: "Default — any available node",
- },
- ...(pools || []).map((pool) => ({
- value: poolId(pool),
- inputDisplay: `${pool.taint.key}=${pool.taint.value || "\"\""}`,
- dropdownDisplay: (
-
- {`${pool.taint.key}=${pool.taint.value || "\"\""}`}
-
- {`taint ${poolId(pool)} · selector ${JSON.stringify(pool.node_selector || {})}`}
-
-
- ),
- })),
- ];
-
- const onChange = (value) => {
- setSelected(value);
- if (value === DEFAULT_VALUE) {
- onSelect({}, []);
- return;
- }
- const pool = (pools || []).find((p) => poolId(p) === value);
- if (!pool) return;
- const toleration = {
- key: pool.taint.key,
- operator: "Equal",
- ...(pool.taint.value ? { value: pool.taint.value } : {}),
- ...(pool.taint.effect ? { effect: pool.taint.effect } : {}),
- };
- onSelect({ ...(pool.node_selector || {}) }, [toleration]);
- };
-
- return (
-
- );
-};
diff --git a/ui/src/pages/version/components/forms/steps/ModelStep.js b/ui/src/pages/version/components/forms/steps/ModelStep.js
index 02b902f04..856dd70a4 100644
--- a/ui/src/pages/version/components/forms/steps/ModelStep.js
+++ b/ui/src/pages/version/components/forms/steps/ModelStep.js
@@ -17,7 +17,7 @@ import { CPULimitsFormGroup } from "../components/CPULimitsFormGroup";
import { ProbesFormGroup } from "../components/ProbesFormGroup";
import { TolerationFormGroup } from "../components/TolerationFormGroup";
import { NodeSelectorFormGroup } from "../components/NodeSelectorFormGroup";
-import { NodePoolSelect } from "../components/NodePoolSelect";
+import { NodeSelect } from "../components/NodeSelect";
export const ModelStep = ({ version, isEnvironmentDisabled = false, maxAllowedReplica, setMaxAllowedReplica }) => {
const { data, onChangeHandler } = useContext(FormContext);
@@ -70,17 +70,18 @@ export const ModelStep = ({ version, isEnvironmentDisabled = false, maxAllowedRe
/>
Node Pool}
- description="Pick a pool to pin this model onto and tolerate — sets the node selector and toleration for you."
+ title={Node
}
+ description="Pick a node to pin this model onto — pins both the predictor and the transformer to the same node, and tolerates the node's taints."
fullWidth
>
- {
+ // Pin predictor and transformer to the same node.
onChange("resource_request.node_selector")(sel);
onChange("resource_request.tolerations")(tols);
+ onChange("transformer.resource_request.node_selector")(sel);
+ onChange("transformer.resource_request.tolerations")(tols);
}}
/>
diff --git a/ui/src/pages/version/components/forms/steps/TransformerStep.js b/ui/src/pages/version/components/forms/steps/TransformerStep.js
index 07c7ca3a6..38a2d20dd 100644
--- a/ui/src/pages/version/components/forms/steps/TransformerStep.js
+++ b/ui/src/pages/version/components/forms/steps/TransformerStep.js
@@ -13,8 +13,6 @@ import { LoggerPanel } from "../components/LoggerPanel";
import { ResourcesPanel } from "../components/ResourcesPanel";
import { SelectTransformerPanel } from "../components/SelectTransformerPanel";
import { CPULimitsFormGroup } from "../components/CPULimitsFormGroup";
-import { TolerationFormGroup } from "../components/TolerationFormGroup";
-import { NodeSelectorFormGroup } from "../components/NodeSelectorFormGroup";
export const TransformerStep = ({ maxAllowedReplica }) => {
const {
@@ -54,17 +52,9 @@ export const TransformerStep = ({ maxAllowedReplica }) => {
onChangeHandler={onChange("transformer.resource_request")}
errors={get(errors, "transformer.resource_request")}
/>
-
-
-
-
+ {/* Node placement (node selector + tolerations) is set once via the
+ Node dropdown in the model step and applied to both the predictor
+ and the transformer, so they always land on the same node. */}
}
/>
From 1c375a8b276401a7fa9c54d4d2ff95e7b3dff99f Mon Sep 17 00:00:00 2001
From: vishwajeetpal
Date: Mon, 20 Jul 2026 14:37:11 +0530
Subject: [PATCH 14/14] added node list and selector
---
.../components/forms/components/NodeSelect.js | 87 +++++++++++++++++++
1 file changed, 87 insertions(+)
create mode 100644 ui/src/pages/version/components/forms/components/NodeSelect.js
diff --git a/ui/src/pages/version/components/forms/components/NodeSelect.js b/ui/src/pages/version/components/forms/components/NodeSelect.js
new file mode 100644
index 000000000..e0cec365a
--- /dev/null
+++ b/ui/src/pages/version/components/forms/components/NodeSelect.js
@@ -0,0 +1,87 @@
+import React, { Fragment, useState } from "react";
+import { EuiSuperSelect, EuiText } from "@elastic/eui";
+import { useMerlinApi } from "../../../../../hooks/useMerlinApi";
+
+const DEFAULT_VALUE = "__default__";
+const HOSTNAME_LABEL = "kubernetes.io/hostname";
+
+const taintSummary = (taints) =>
+ taints && taints.length > 0
+ ? taints
+ .map((t) => `${t.key}=${t.value || ""}:${t.effect || ""}`)
+ .join(", ")
+ : "no taints";
+
+/**
+ * NodeSelect - a dropdown to pin a model onto a specific node.
+ *
+ * Selecting a node sets the node selector to that node's hostname (pin) and adds
+ * a toleration for each of the node's taints (permit), so the pod can land there.
+ *
+ * Props
+ * environment – environment name (used to fetch that cluster's nodes)
+ * onSelect(sel, tols) – called with the new node_selector + tolerations
+ */
+export const NodeSelect = ({ environment, onSelect }) => {
+ const [selected, setSelected] = useState(DEFAULT_VALUE);
+
+ const [{ data: nodes, isLoaded }] = useMerlinApi(
+ `/environments/${environment}/nodes`,
+ {},
+ [],
+ !!environment
+ );
+
+ const options = [
+ {
+ value: DEFAULT_VALUE,
+ inputDisplay: "Default — any available node",
+ },
+ ...(nodes || []).map((node) => ({
+ value: node.name,
+ inputDisplay: node.name,
+ disabled: !node.ready,
+ dropdownDisplay: (
+
+
+ {node.name}
+ {!node.ready ? " (not ready)" : ""}
+
+
+ {taintSummary(node.taints)}
+
+
+ ),
+ })),
+ ];
+
+ const onChange = (value) => {
+ setSelected(value);
+ if (value === DEFAULT_VALUE) {
+ onSelect({}, []);
+ return;
+ }
+ const node = (nodes || []).find((n) => n.name === value);
+ if (!node) return;
+
+ const nodeSelector = { [HOSTNAME_LABEL]: node.name };
+ const tolerations = (node.taints || []).map((t) => ({
+ key: t.key,
+ operator: t.value ? "Equal" : "Exists",
+ ...(t.value ? { value: t.value } : {}),
+ ...(t.effect ? { effect: t.effect } : {}),
+ }));
+ onSelect(nodeSelector, tolerations);
+ };
+
+ return (
+
+ );
+};