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 (
{taintSummary(node.taints)}