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/api/api/node_api.go b/api/api/node_api.go new file mode 100644 index 000000000..253ebe3a7 --- /dev/null +++ b/api/api/node_api.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 api + +import ( + "errors" + "fmt" + "net/http" + + "gorm.io/gorm" +) + +// NodeController serves node listing for a deployment environment. +type NodeController struct { + *AppContext +} + +// 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"] + 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)) + } + + // 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). + nodes, err := c.NodeService.ListNodes(ctx, env.Name) + if err != nil { + return InternalServerError(fmt.Sprintf("Error listing nodes: %v", err)) + } + + return Ok(nodes) +} diff --git a/api/api/router.go b/api/api/router.go index 6fc8902d8..c548f734e 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 + NodeService service.NodeService 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} + nodeController := NodeController{&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}/nodes", nil, nodeController.ListNodes, "ListNodes"}, // 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 b75cbc17a..ea361dedd 100644 --- a/api/api/validator.go +++ b/api/api/validator.go @@ -4,8 +4,10 @@ import ( "context" "errors" "fmt" + "strings" "golang.org/x/exp/slices" + corev1 "k8s.io/api/core/v1" "github.com/caraml-dev/merlin/config" "github.com/caraml-dev/merlin/models" @@ -61,22 +63,86 @@ func validateRequest(validators ...requestValidator) error { func resourceRequestValidation(endpoint *models.VersionEndpoint) requestValidator { return newFuncValidate(func() error { - if endpoint.ResourceRequest == nil { - return nil - } + 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.MaxReplica < 1 { + 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) + } - if endpoint.ResourceRequest.MinReplica > endpoint.ResourceRequest.MaxReplica { - return fmt.Errorf("min replica must be less or equal to max replica") + if err := validateNodeSelector(endpoint.ResourceRequest.NodeSelector); err != nil { + return fmt.Errorf("invalid node selector 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.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 }) } +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 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/api/validator_test.go b/api/api/validator_test.go new file mode 100644 index 000000000..65085e2d4 --- /dev/null +++ b/api/api/validator_test.go @@ -0,0 +1,276 @@ +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 + trans *models.Transformer + 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: "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{ + MinReplica: 1, MaxReplica: 2, + Tolerations: []corev1.Toleration{ + {Key: "k", Operator: "BadOp", Value: "v"}, + }, + }, + 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, + Transformer: tt.trans, + } + validator := resourceRequestValidation(endpoint) + err := validator.validate() + if tt.wantErr { + assert.Error(t, err) + } else { + assert.NoError(t, err) + } + }) + } +} + +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/client/model_resource_request.go b/api/client/model_resource_request.go index ccd65559a..de104ca8f 100644 --- a/api/client/model_resource_request.go +++ b/api/client/model_resource_request.go @@ -19,13 +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"` + 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 @@ -269,6 +271,68 @@ 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 +} + +// 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 { @@ -300,6 +364,12 @@ 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 + } + if !IsNil(o.NodeSelector) { + toSerialize["node_selector"] = o.NodeSelector + } 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/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 3018e1bc6..2b3336242 100644 --- a/api/cluster/resource/templater.go +++ b/api/cluster/resource/templater.go @@ -230,13 +230,28 @@ func (t *InferenceServiceTemplater) createPredictorSpec(modelService *models.Ser resources.Requests[resourceType] = resourceQuantity resources.Limits[resourceType] = resourceQuantity - nodeSelector = gpuConfig.NodeSelector - tolerations = gpuConfig.Tolerations + // 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...) } } } } + // 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...) + } + + // 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 @@ -468,6 +483,16 @@ func (t *InferenceServiceTemplater) createTransformerSpec( }, } + // Apply user-defined tolerations for transformer pods + if len(transformer.ResourceRequest.Tolerations) > 0 { + 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 ff0e9e620..e608f0c2d 100644 --- a/api/cluster/resource/templater_test.go +++ b/api/cluster/resource/templater_test.go @@ -4794,3 +4794,328 @@ 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) + }) + } +} + +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..072b34956 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) + 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) @@ -352,6 +353,7 @@ func buildDependencies(ctx context.Context, cfg *config.Config, db *gorm.DB, dis VersionImageService: versionImageService, EndpointsService: versionEndpointService, LogService: logService, + 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_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/models/resource_request.go b/api/models/resource_request.go index d335c9632..d370908da 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,10 @@ 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"` + // 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_service.go b/api/service/node_service.go new file mode 100644 index 000000000..42e645c97 --- /dev/null +++ b/api/service/node_service.go @@ -0,0 +1,57 @@ +// 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" +) + +// 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 nodeService struct { + // clusterControllers is a map of environment name to its cluster.Controller. + clusterControllers map[string]cluster.Controller +} + +// NewNodeService creates a NodeService. +// clusterControllers is a map of environment name to its cluster.Controller. +func NewNodeService(clusterControllers map[string]cluster.Controller) NodeService { + return &nodeService{clusterControllers: clusterControllers} +} + +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) + } + + nodeList, err := controller.ListNodes(ctx) + if err != nil { + return nil, fmt.Errorf("unable to list nodes for environment %s: %w", environmentName, err) + } + + return models.NewNodeInfos(nodeList.Items), nil +} 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)) -} 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 c92773f23..df4c55ebb 100644 --- a/scripts/e2e/config/istio/ingress-gateway.yaml +++ b/scripts/e2e/config/istio/ingress-gateway.yaml @@ -5,3 +5,8 @@ resources: requests: cpu: 50m memory: 64Mi + +securityContext: + runAsUser: 1337 + runAsGroup: 1337 + runAsNonRoot: true \ No newline at end of file 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 diff --git a/swagger.yaml b/swagger.yaml index 65bece648..f66d40b89 100644 --- a/swagger.yaml +++ b/swagger.yaml @@ -67,6 +67,27 @@ paths: type: array items: "$ref": "#/components/schemas/Environment" + "/environments/{environment_name}/nodes": + get: + tags: + - environment + 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 + required: true + schema: + type: string + responses: + "200": + description: OK + content: + "*/*": + schema: + type: array + items: + "$ref": "#/components/schemas/NodeInfo" "/projects": get: tags: @@ -2121,6 +2142,61 @@ 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" + node_selector: + type: object + description: NodeSelector pins model pods onto nodes whose labels match every entry + additionalProperties: + type: string + NodeInfo: + type: object + description: A cluster node's placement-relevant summary + properties: + name: + type: string + ready: + type: boolean + labels: + type: object + 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 + 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: diff --git a/ui/src/components/ResourcesConfigTable.js b/ui/src/components/ResourcesConfigTable.js index dcf9a95cd..6dcee7db5 100644 --- a/ui/src/components/ResourcesConfigTable.js +++ b/ui/src/components/ResourcesConfigTable.js @@ -29,6 +29,8 @@ export const ResourcesConfigTable = ({ gpu_request, liveness_probe, readiness_probe, + tolerations, + node_selector, }, }) => { const items = [ @@ -104,6 +106,31 @@ 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(", ") || "—", + }); + }); + } + + // 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 ( 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; + }); + } + } + // Drop empty node selectors so the backend doesn't receive an empty object + if (_.isEmpty(versionEndpoint?.resource_request?.node_selector)) { + delete versionEndpoint?.resource_request?.node_selector; + } + if (_.isEmpty(versionEndpoint?.transformer?.resource_request?.node_selector)) { + delete versionEndpoint?.transformer?.resource_request?.node_selector; + } submitForm({ body: JSON.stringify({ ...versionEndpoint, 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 ( + + ); +}; 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/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..856dd70a4 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"; @@ -15,6 +15,9 @@ 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"; +import { NodeSelectorFormGroup } from "../components/NodeSelectorFormGroup"; +import { NodeSelect } from "../components/NodeSelect"; export const ModelStep = ({ version, isEnvironmentDisabled = false, maxAllowedReplica, setMaxAllowedReplica }) => { const { data, onChangeHandler } = useContext(FormContext); @@ -65,6 +68,34 @@ export const ModelStep = ({ version, isEnvironmentDisabled = false, maxAllowedRe onChangeHandler={onChange("image_builder_resource_request")} errors={get(errors, "image_builder_resource_request")} /> + + 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 f2f6fe7ac..38a2d20dd 100644 --- a/ui/src/pages/version/components/forms/steps/TransformerStep.js +++ b/ui/src/pages/version/components/forms/steps/TransformerStep.js @@ -52,6 +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. */} } /> diff --git a/ui/src/services/transformer/Transformer.js b/ui/src/services/transformer/Transformer.js index 01bd4fb5b..c90745b30 100644 --- a/ui/src/services/transformer/Transformer.js +++ b/ui/src/services/transformer/Transformer.js @@ -23,7 +23,9 @@ 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: [], + node_selector: {}, }; this.env_vars = []; @@ -60,6 +62,14 @@ export class Transformer { transformer.secrets = []; } + if (transformer.resource_request && !transformer.resource_request.tolerations) { + 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 465058ddc..3b45a5e5a 100644 --- a/ui/src/services/version_endpoint/VersionEndpoint.js +++ b/ui/src/services/version_endpoint/VersionEndpoint.js @@ -30,6 +30,8 @@ export class VersionEndpoint { memory_request: "512Mi", liveness_probe: null, readiness_probe: null, + tolerations: [], + node_selector: {}, }; this.image_builder_resource_request = { @@ -74,6 +76,14 @@ export class VersionEndpoint { } } + if (versionEndpoint.resource_request && !versionEndpoint.resource_request.tolerations) { + 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); }