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 ( + + ); +};