diff --git a/internal/entity/models/base_model.go b/internal/entity/models/base_model.go index c0b0d2d4fb..2662adb607 100644 --- a/internal/entity/models/base_model.go +++ b/internal/entity/models/base_model.go @@ -400,6 +400,7 @@ func ParseListModel(modelList ModelList) []ListModelResponse { modelResponse.Name = modelName if modelEntity != nil { modelResponse.MaxDimension = modelEntity.MaxDimension + modelResponse.MaxBatchSize = modelEntity.MaxBatchSize modelResponse.Dimensions = modelEntity.Dimensions modelResponse.MaxOutput = modelEntity.MaxOutput modelResponse.ModelTypes = modelEntity.ModelTypes diff --git a/internal/entity/models/model.go b/internal/entity/models/model.go index 17069d9095..330e0aa16e 100644 --- a/internal/entity/models/model.go +++ b/internal/entity/models/model.go @@ -166,7 +166,8 @@ type Model struct { Thinking *ModelThinking `json:"thinking"` Tools *ModelTools `json:"tools"` Class *string `json:"class"` - MaxDimension *int `json:"max_dimension"` // used by embedding models + MaxDimension *int `json:"max_dimension"` // used by embedding models + MaxBatchSize *int `json:"max_batch_size"` // used by embedding models Dimensions []int `json:"dimensions"` BatchSize *int `json:"batch_size"` // max texts per Embed request; used by embedding models Alias []string `json:"alias"` @@ -417,6 +418,9 @@ func (pm *ProviderManager) ListAllModels() ([]map[string]interface{}, error) { if model.MaxDimension != nil { modelData["max_dimension"] = *model.MaxDimension } + if model.MaxBatchSize != nil { + modelData["max_batch_size"] = *model.MaxBatchSize + } if len(model.Dimensions) > 0 { modelData["dimensions"] = model.Dimensions } @@ -516,11 +520,12 @@ func (pm *ProviderManager) ListModels(providerName string) ([]map[string]interfa // keep the response shape stable for clients that destructure // the object. modelData := map[string]interface{}{ - "name": model.Name, - "max_output": model.MaxOutput, - "model_types": model.ModelTypes, - "max_dimension": model.MaxDimension, - "dimensions": model.Dimensions, + "name": model.Name, + "max_output": model.MaxOutput, + "model_types": model.ModelTypes, + "max_dimension": model.MaxDimension, + "max_batch_size": model.MaxBatchSize, + "dimensions": model.Dimensions, } if model.BatchSize != nil { modelData["batch_size"] = *model.BatchSize @@ -634,9 +639,10 @@ func (pm *ProviderManager) SearchModelInfo(providerName, modelName string, filte if matchFilter { modelData := map[string]interface{}{ - "name": model.Name, - "max_output": model.MaxOutput, - "model_types": model.ModelTypes, + "name": model.Name, + "max_output": model.MaxOutput, + "model_types": model.ModelTypes, + "max_batch_size": model.MaxBatchSize, //"features": getFeaturesMap(model.Features), } @@ -694,10 +700,11 @@ func (pm *ProviderManager) SearchByType(modelType string) ModelResponse { for _, model := range provider.Models { if containsModelType(model.ModelTypes, modelType) { modelData := map[string]interface{}{ - "provider": provider.Name, - "name": model.Name, - "max_output": model.MaxOutput, - "model_types": model.ModelTypes, + "provider": provider.Name, + "name": model.Name, + "max_output": model.MaxOutput, + "model_types": model.ModelTypes, + "max_batch_size": model.MaxBatchSize, //"features": getFeaturesMap(model.Features), } resp.Data = append(resp.Data, modelData) diff --git a/internal/entity/models/nvidia.go b/internal/entity/models/nvidia.go index 586090b30e..c8e6a00665 100644 --- a/internal/entity/models/nvidia.go +++ b/internal/entity/models/nvidia.go @@ -694,6 +694,7 @@ func parseNvidiaModelList(modelList ModelList, provider *Provider) []ListModelRe response.ModelTypes = append([]string(nil), preset.ModelTypes...) response.Thinking = preset.Thinking response.MaxDimension = preset.MaxDimension + response.MaxBatchSize = preset.MaxBatchSize response.Dimensions = append([]int(nil), preset.Dimensions...) } else { maxTokens := defaultMaxTokens diff --git a/internal/entity/models/types.go b/internal/entity/models/types.go index 1305c8c72f..e284f8247d 100644 --- a/internal/entity/models/types.go +++ b/internal/entity/models/types.go @@ -110,7 +110,8 @@ type ListModelResponse struct { MaxOutput *int `json:"max_output"` ModelTypes []string `json:"model_types"` Thinking *ModelThinking `json:"thinking"` - MaxDimension *int `json:"max_dimension"` // used by embedding models + MaxDimension *int `json:"max_dimension"` // used by embedding models + MaxBatchSize *int `json:"max_batch_size"` // used by embedding models Dimensions []int `json:"dimensions"` } diff --git a/internal/service/model_service.go b/internal/service/model_service.go index 4ff4b0ed0c..483b902a12 100644 --- a/internal/service/model_service.go +++ b/internal/service/model_service.go @@ -398,6 +398,7 @@ func (m *ModelProviderService) ListSupportedModels(ctx context.Context, provider result = append(result, map[string]interface{}{ "name": model.Name, "max_dimension": model.MaxDimension, + "max_batch_size": model.MaxBatchSize, "dimensions": model.Dimensions, "content_length": model.ContentLength, "max_output": model.MaxOutput, @@ -515,6 +516,9 @@ func setDiscoveredModelMetadata(extra map[string]interface{}, model modelModule. if model.MaxDimension != nil { extra["max_dimension"] = *model.MaxDimension } + if model.MaxBatchSize != nil { + extra["max_batch_size"] = *model.MaxBatchSize + } if len(model.Dimensions) > 0 { extra["dimensions"] = model.Dimensions } @@ -1140,8 +1144,11 @@ func verifyProviderModel(ctx context.Context, driver modelModule.ModelDriver, pr modelTypes = modelModule.InferModelTypes(modelName) } modelsToVerify = append(modelsToVerify, &modelModule.Model{ - Name: modelName, - ModelTypes: modelTypes, + Name: modelName, + ModelTypes: modelTypes, + MaxDimension: rm.MaxDimension, + MaxBatchSize: rm.MaxBatchSize, + Dimensions: append([]int(nil), rm.Dimensions...), }) } } else { @@ -1178,7 +1185,19 @@ func verifyProviderModel(ctx context.Context, driver modelModule.ModelDriver, pr msg := []modelModule.Message{{Role: "user", Content: "Hi"}} _, err = driver.ChatWithMessages(ctx, modelName, msg, apiConfig, nil, nil) case "embedding": - _, err = driver.Embed(ctx, &modelName, []string{"test"}, apiConfig, nil, nil) + // Provider discovery can return models without catalog limits. Apply + // the strict validator whenever the model has the metadata required + // to construct a valid verification request. + if model.MaxDimension != nil && model.MaxBatchSize != nil { + requestedDimension := *model.MaxDimension + if len(model.Dimensions) > 0 { + requestedDimension = model.Dimensions[0] + } + err = validateEmbeddingModel(model, requestedDimension, 1) + } + if err == nil { + _, err = driver.Embed(ctx, &modelName, []string{"test"}, apiConfig, nil, nil) + } case "rerank": _, err = driver.Rerank(ctx, &modelName, "test", []string{"test"}, apiConfig, &modelModule.RerankConfig{}, nil) case "tts": @@ -2509,6 +2528,7 @@ type tenantModelExtra struct { MaxTokens *int `json:"max_tokens"` ModelTypes []string `json:"model_types"` MaxDimension *int `json:"max_dimension"` + MaxBatchSize *int `json:"max_batch_size"` Dimensions []int `json:"dimensions"` Thinking *bool `json:"thinking"` } @@ -2551,6 +2571,9 @@ func modelInfoWithTenantExtra(modelInfo *modelModule.Model, modelEntity *entity. if extra.MaxDimension != nil && *extra.MaxDimension > 0 { model.MaxDimension = extra.MaxDimension } + if extra.MaxBatchSize != nil && *extra.MaxBatchSize > 0 { + model.MaxBatchSize = extra.MaxBatchSize + } if len(extra.Dimensions) > 0 { model.Dimensions = append([]int(nil), extra.Dimensions...) } @@ -2860,28 +2883,48 @@ func (m *ModelProviderService) ChatToModelStreamWithSender(ctx context.Context, return common.CodeSuccess, nil } -func validateEmbeddingDimension(model *modelModule.Model, requested int) error { - if requested <= 0 || model == nil { - return nil +func validateEmbeddingModel(model *modelModule.Model, requestedDimension, requestedBatchSize int) error { + if model == nil { + return fmt.Errorf("embedding model is nil") + } + + if requestedDimension <= 0 { + return fmt.Errorf("input dimension <= 0") + } + + if requestedBatchSize <= 0 { + return fmt.Errorf("input batch size <= 0") + } + + if model.MaxDimension == nil { + return fmt.Errorf("input embedding max dimension is nil, %s", model.Name) + } + + if model.MaxBatchSize == nil { + return fmt.Errorf("input embedding max batch size is nil, %s", model.Name) + } + + if *model.MaxBatchSize < requestedBatchSize { + return fmt.Errorf("input embedding max batch size is more than limitation, %s", model.Name) } if len(model.Dimensions) > 0 { for _, dim := range model.Dimensions { - if dim == requested { + if dim == requestedDimension { return nil } } return fmt.Errorf( "dimension %d is not supported by model %s, supported dimensions: %v", - requested, + requestedDimension, model.Name, model.Dimensions, ) } - if model.MaxDimension != nil && requested > *model.MaxDimension { + if model.MaxDimension != nil && requestedDimension > *model.MaxDimension { return fmt.Errorf( "dimension %d is not supported by model %s, max dimension: %d", - requested, + requestedDimension, model.Name, *model.MaxDimension, ) @@ -2940,7 +2983,7 @@ func (m *ModelProviderService) EmbedText(ctx context.Context, providerName, inst } } - if err = validateEmbeddingDimension(info.ModelInfo, modelConfig.Dimension); err != nil { + if err = validateEmbeddingModel(info.ModelInfo, modelConfig.Dimension, len(texts)); err != nil { return nil, common.CodeBadRequest, err } diff --git a/internal/service/model_service_test.go b/internal/service/model_service_test.go index 8059a63772..86cc46cc3d 100644 --- a/internal/service/model_service_test.go +++ b/internal/service/model_service_test.go @@ -15,79 +15,185 @@ import ( modelModule "ragflow/internal/entity/models" ) -func TestValidateEmbeddingDimension(t *testing.T) { +type remoteModelProbeDriver struct { + *modelModule.DummyModel + remoteModels []modelModule.ListModelResponse + embedCalls int +} + +func (d *remoteModelProbeDriver) ListModels(context.Context, *modelModule.APIConfig) ([]modelModule.ListModelResponse, error) { + return d.remoteModels, nil +} + +func (d *remoteModelProbeDriver) Embed(context.Context, *string, []string, *modelModule.APIConfig, *modelModule.EmbeddingConfig, *common.ModelUsage) ([]modelModule.EmbeddingData, error) { + d.embedCalls++ + return nil, nil +} + +func TestValidateEmbeddingModel(t *testing.T) { maxDimension := 2048 + maxBatchSize := 128 tests := []struct { - name string - model *modelModule.Model - requested int - wantErr string + name string + model *modelModule.Model + requestedDimension int + requestedBatchSize int + wantErr string }{ { - name: "allows unset requested dimension", - model: &modelModule.Model{MaxDimension: &maxDimension, Dimensions: []int{256, 512}}, - requested: 0, + name: "rejects nil model", + requestedDimension: 1024, + requestedBatchSize: 16, + wantErr: "embedding model is nil", }, { - name: "allows missing model schema", - model: nil, - requested: 256, + name: "rejects zero dimension", + model: &modelModule.Model{}, + requestedDimension: 0, + requestedBatchSize: 1, + wantErr: "input dimension <= 0", }, { - name: "allows dimension listed in explicit options", - model: &modelModule.Model{Name: "embedding-3", MaxDimension: &maxDimension, Dimensions: []int{256, 512, 1024, 2048}}, - requested: 1024, + name: "rejects negative dimension", + model: &modelModule.Model{}, + requestedDimension: -1, + requestedBatchSize: 1, + wantErr: "input dimension <= 0", }, { - name: "rejects dimension not listed in explicit options", - model: &modelModule.Model{Name: "embedding-3", MaxDimension: &maxDimension, Dimensions: []int{256, 512, 1024, 2048}}, - requested: 1536, - wantErr: "supported dimensions", + name: "rejects zero batch size", + model: &modelModule.Model{}, + requestedDimension: 1024, + requestedBatchSize: 0, + wantErr: "input batch size <= 0", }, { - name: "allows custom dimension within max dimension", - model: &modelModule.Model{Name: "flex-embedding", MaxDimension: &maxDimension}, - requested: 1536, + name: "rejects negative batch size", + model: &modelModule.Model{}, + requestedDimension: 1024, + requestedBatchSize: -1, + wantErr: "input batch size <= 0", }, { - name: "rejects custom dimension above max dimension", - model: &modelModule.Model{Name: "flex-embedding", MaxDimension: &maxDimension}, - requested: 4096, - wantErr: "max dimension", + name: "rejects missing max dimension", + model: &modelModule.Model{MaxBatchSize: &maxBatchSize}, + requestedDimension: 1024, + requestedBatchSize: 1, + wantErr: "max dimension is nil", + }, + { + name: "rejects missing max batch size", + model: &modelModule.Model{MaxDimension: &maxDimension}, + requestedDimension: 1024, + requestedBatchSize: 1, + wantErr: "max batch size is nil", + }, + { + name: "allows dimension listed in explicit options", + model: &modelModule.Model{Name: "embedding-3", MaxDimension: &maxDimension, MaxBatchSize: &maxBatchSize, Dimensions: []int{256, 512, 1024, 2048}}, + requestedDimension: 1024, + requestedBatchSize: 128, + }, + { + name: "rejects dimension not listed in explicit options", + model: &modelModule.Model{Name: "embedding-3", MaxDimension: &maxDimension, MaxBatchSize: &maxBatchSize, Dimensions: []int{256, 512, 1024, 2048}}, + requestedDimension: 1536, + requestedBatchSize: 128, + wantErr: "supported dimensions", + }, + { + name: "allows custom dimension within max dimension", + model: &modelModule.Model{Name: "flex-embedding", MaxDimension: &maxDimension, MaxBatchSize: &maxBatchSize}, + requestedDimension: 1536, + requestedBatchSize: 1, + }, + { + name: "rejects custom dimension above max dimension", + model: &modelModule.Model{Name: "flex-embedding", MaxDimension: &maxDimension, MaxBatchSize: &maxBatchSize}, + requestedDimension: 4096, + requestedBatchSize: 1, + wantErr: "max dimension", + }, + { + name: "allows batch at model limit", + model: &modelModule.Model{Name: "embedding-3", MaxDimension: &maxDimension, MaxBatchSize: &maxBatchSize}, + requestedDimension: 1024, + requestedBatchSize: 128, + }, + { + name: "rejects batch above model limit", + model: &modelModule.Model{Name: "embedding-3", MaxDimension: &maxDimension, MaxBatchSize: &maxBatchSize}, + requestedDimension: 1024, + requestedBatchSize: 129, + wantErr: "max batch size", + }, + { + name: "rejects batch when model limit is unspecified", + model: &modelModule.Model{Name: "custom-embedding", MaxDimension: &maxDimension}, + requestedDimension: 1024, + requestedBatchSize: 10000, + wantErr: "max batch size is nil", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - err := validateEmbeddingDimension(tt.model, tt.requested) + err := validateEmbeddingModel(tt.model, tt.requestedDimension, tt.requestedBatchSize) if tt.wantErr == "" { if err != nil { - t.Fatalf("validateEmbeddingDimension() error = %v", err) + t.Fatalf("validateEmbeddingModel() error = %v", err) } return } if err == nil { - t.Fatalf("validateEmbeddingDimension() expected error containing %q", tt.wantErr) + t.Fatalf("validateEmbeddingModel() expected error containing %q", tt.wantErr) } if !strings.Contains(err.Error(), tt.wantErr) { - t.Fatalf("validateEmbeddingDimension() error = %v, want substring %q", err, tt.wantErr) + t.Fatalf("validateEmbeddingModel() error = %v, want substring %q", err, tt.wantErr) } }) } } -func TestModelInfoWithTenantExtraAppliesEmbeddingDimensions(t *testing.T) { +func TestVerifyProviderModelValidatesRemoteEmbeddingMetadata(t *testing.T) { + maxDimension := 1024 + maxBatchSize := 0 + driver := &remoteModelProbeDriver{ + DummyModel: modelModule.NewDummyModel(nil, modelModule.URLSuffix{}), + remoteModels: []modelModule.ListModelResponse{{ + Name: "remote-embedding", + ModelTypes: []string{"embedding"}, + MaxDimension: &maxDimension, + MaxBatchSize: &maxBatchSize, + }}, + } + + result, err := verifyProviderModel(context.Background(), driver, nil, &modelModule.APIConfig{}, nil) + if err == nil { + t.Fatal("verifyProviderModel() error = nil, want validation error") + } + if result["remote-embedding"] != entity.ModelVerifyFail { + t.Fatalf("verification result = %#v, want remote model failure", result) + } + if driver.embedCalls != 0 { + t.Fatalf("Embed calls = %d, want 0 after metadata validation failure", driver.embedCalls) + } +} + +func TestModelInfoWithTenantExtraAppliesEmbeddingConstraints(t *testing.T) { factoryMaxDimension := 2048 + factoryBatchSize := 128 modelInfo := &modelModule.Model{ Name: "embedding-3", MaxDimension: &factoryMaxDimension, + MaxBatchSize: &factoryBatchSize, Dimensions: []int{1024, 2048}, ModelTypes: []string{"embedding"}, ModelTypeMap: map[string]bool{"embedding": true}, } modelEntity := &entity.TenantModel{ - Extra: `{"max_dimension":768,"dimensions":[384,768],"model_types":["embedding"]}`, + Extra: `{"max_dimension":768,"max_batch_size":16,"dimensions":[384,768],"model_types":["embedding"]}`, } merged, err := modelInfoWithTenantExtra(modelInfo, modelEntity) @@ -100,18 +206,27 @@ func TestModelInfoWithTenantExtraAppliesEmbeddingDimensions(t *testing.T) { if merged.MaxDimension == nil || *merged.MaxDimension != 768 { t.Fatalf("MaxDimension = %v, want 768", merged.MaxDimension) } + if merged.MaxBatchSize == nil || *merged.MaxBatchSize != 16 { + t.Fatalf("MaxBatchSize = %v, want 16", merged.MaxBatchSize) + } if len(merged.Dimensions) != 2 || merged.Dimensions[0] != 384 || merged.Dimensions[1] != 768 { t.Fatalf("Dimensions = %v, want [384 768]", merged.Dimensions) } - if err := validateEmbeddingDimension(merged, 1024); err == nil || !strings.Contains(err.Error(), "supported dimensions") { - t.Fatalf("validateEmbeddingDimension() error = %v, want supported dimensions error", err) + if validationErr := validateEmbeddingModel(merged, 1024, 16); validationErr == nil || !strings.Contains(validationErr.Error(), "supported dimensions") { + t.Fatalf("validateEmbeddingModel() error = %v, want supported dimensions error", validationErr) } - if err := validateEmbeddingDimension(merged, 768); err != nil { - t.Fatalf("validateEmbeddingDimension() error = %v", err) + if validationErr := validateEmbeddingModel(merged, 768, 16); validationErr != nil { + t.Fatalf("validateEmbeddingModel() error = %v", validationErr) + } + if validationErr := validateEmbeddingModel(merged, 768, 17); validationErr == nil || !strings.Contains(validationErr.Error(), "max batch size") { + t.Fatalf("validateEmbeddingModel() error = %v, want max batch size error", validationErr) } if modelInfo.MaxDimension == nil || *modelInfo.MaxDimension != factoryMaxDimension { t.Fatalf("factory MaxDimension was mutated: %v", modelInfo.MaxDimension) } + if modelInfo.MaxBatchSize == nil || *modelInfo.MaxBatchSize != factoryBatchSize { + t.Fatalf("factory MaxBatchSize was mutated: %v", modelInfo.MaxBatchSize) + } if len(modelInfo.Dimensions) != 2 || modelInfo.Dimensions[0] != 1024 || modelInfo.Dimensions[1] != 2048 { t.Fatalf("factory Dimensions were mutated: %v", modelInfo.Dimensions) }