GO: improve model info parsing and add model_id/tenant context to list response (#16804)

This commit is contained in:
buua436
2026-07-10 13:29:01 +08:00
committed by GitHub
parent 868e524f29
commit 28340f6218
17 changed files with 661 additions and 159 deletions

View File

@@ -54,17 +54,24 @@ type TenantWithRole struct {
// TenantInfo tenant information with role (for owner tenant)
type TenantInfo struct {
TenantID string `gorm:"column:tenant_id" json:"tenant_id"`
Name *string `gorm:"column:name" json:"name,omitempty"`
LLMID string `gorm:"column:llm_id" json:"llm_id"`
EmbDID string `gorm:"column:embd_id" json:"embd_id"`
RerankID string `gorm:"column:rerank_id" json:"rerank_id"`
ASRID string `gorm:"column:asr_id" json:"asr_id"`
Img2TxtID string `gorm:"column:img2txt_id" json:"img2txt_id"`
TTSID *string `gorm:"column:tts_id" json:"tts_id,omitempty"`
OCRID *string `gorm:"column:ocr_id" json:"ocr_id,omitempty"`
ParserIDs string `gorm:"column:parser_ids" json:"parser_ids"`
Role string `gorm:"column:role" json:"role"`
TenantID string `gorm:"column:tenant_id" json:"tenant_id"`
Name *string `gorm:"column:name" json:"name,omitempty"`
LLMID string `gorm:"column:llm_id" json:"llm_id"`
TenantLLMID *string `gorm:"column:tenant_llm_id" json:"tenant_llm_id,omitempty"`
EmbDID string `gorm:"column:embd_id" json:"embd_id"`
TenantEmbdID *string `gorm:"column:tenant_embd_id" json:"tenant_embd_id,omitempty"`
RerankID string `gorm:"column:rerank_id" json:"rerank_id"`
TenantRerankID *string `gorm:"column:tenant_rerank_id" json:"tenant_rerank_id,omitempty"`
ASRID string `gorm:"column:asr_id" json:"asr_id"`
TenantASRID *string `gorm:"column:tenant_asr_id" json:"tenant_asr_id,omitempty"`
Img2TxtID string `gorm:"column:img2txt_id" json:"img2txt_id"`
TenantImg2TxtID *string `gorm:"column:tenant_img2txt_id" json:"tenant_img2txt_id,omitempty"`
TTSID *string `gorm:"column:tts_id" json:"tts_id,omitempty"`
TenantTTSID *string `gorm:"column:tenant_tts_id" json:"tenant_tts_id,omitempty"`
OCRID *string `gorm:"column:ocr_id" json:"ocr_id,omitempty"`
TenantOCRID *string `gorm:"column:tenant_ocr_id" json:"tenant_ocr_id,omitempty"`
ParserIDs string `gorm:"column:parser_ids" json:"parser_ids"`
Role string `gorm:"column:role" json:"role"`
}
// GetInfoByUserID get tenant information for the owner tenant of a user
@@ -72,7 +79,7 @@ func (dao *TenantDAO) GetInfoByUserID(userID string) ([]*TenantInfo, error) {
var results []*TenantInfo
err := DB.Model(&entity.Tenant{}).
Select("tenant.id as tenant_id, tenant.name, tenant.llm_id, tenant.embd_id, tenant.rerank_id, tenant.asr_id, tenant.img2txt_id, tenant.tts_id, tenant.ocr_id, tenant.parser_ids, user_tenant.role").
Select("tenant.id as tenant_id, tenant.name, tenant.llm_id, tenant.tenant_llm_id, tenant.embd_id, tenant.tenant_embd_id, tenant.rerank_id, tenant.tenant_rerank_id, tenant.asr_id, tenant.tenant_asr_id, tenant.img2txt_id, tenant.tenant_img2txt_id, tenant.tts_id, tenant.tenant_tts_id, tenant.ocr_id, tenant.tenant_ocr_id, tenant.parser_ids, user_tenant.role").
Joins("INNER JOIN user_tenant ON user_tenant.tenant_id = tenant.id").
Where("user_tenant.user_id = ? AND user_tenant.status = ? AND user_tenant.role = ? AND tenant.status = ?", userID, "1", "owner", "1").
Scan(&results).Error

View File

@@ -119,6 +119,8 @@ func ModelTypeFromString(s string) ModelType {
return ModelTypeSpeech2Text
case "image2text":
return ModelTypeImage2Text
case "vision":
return ModelTypeImage2Text
case "rerank":
return ModelTypeRerank
case "tts":

View File

@@ -1282,8 +1282,9 @@ func (h *ProviderHandler) ListTenantAddedModels(c *gin.Context) {
}
modelType := c.Query("type")
ownerTenantID := c.Query("owner_tenant_id")
addedModels, code, err := h.modelProviderService.ListTenantAddedModels(user.ID, modelType)
addedModels, code, err := h.modelProviderService.ListTenantAddedModels(user.ID, ownerTenantID, modelType)
if err != nil {
common.ErrorWithCode(c, int(code), err.Error())
return

View File

@@ -201,7 +201,7 @@ func defaultPDFVisionModelResolver(
driver, modelName, apiConfig, _, err := resolveTenantModelByType(tenantID, entity.ModelTypeImage2Text)
return driver, modelName, apiConfig, err
}
driver, modelName, apiConfig, _, err := resolveModelConfigFromProviderInstance(tenantID, entity.ModelTypeImage2Text, modelID)
driver, modelName, apiConfig, _, err := resolveModelConfig(tenantID, entity.ModelTypeImage2Text, modelID)
return driver, modelName, apiConfig, err
}
@@ -303,7 +303,112 @@ func resolveTenantModelByType(tenantID string, modelType entity.ModelType) (mode
if modelID == "" {
return nil, "", nil, 0, fmt.Errorf("no default %s model is set", modelType)
}
return resolveModelConfigFromProviderInstance(tenantID, modelType, modelID)
if tenantModelID := tenantModelIDByType(tenant, modelType); tenantModelID != "" {
driver, modelName, apiConfig, maxTokens, err := resolveModelConfigByID(tenantID, modelType, tenantModelID)
if err == nil {
return driver, modelName, apiConfig, maxTokens, nil
}
}
return resolveModelConfig(tenantID, modelType, modelID)
}
func tenantModelIDByType(tenant *entity.Tenant, modelType entity.ModelType) string {
if tenant == nil {
return ""
}
switch modelType {
case entity.ModelTypeChat:
return stringValue(tenant.TenantLLMID)
case entity.ModelTypeEmbedding:
return stringValue(tenant.TenantEmbdID)
case entity.ModelTypeRerank:
return stringValue(tenant.TenantRerankID)
case entity.ModelTypeSpeech2Text:
return stringValue(tenant.TenantASRID)
case entity.ModelTypeImage2Text:
return stringValue(tenant.TenantImg2TxtID)
case entity.ModelTypeTTS:
return stringValue(tenant.TenantTTSID)
case entity.ModelTypeOCR:
return stringValue(tenant.TenantOCRID)
default:
return ""
}
}
func stringValue(value *string) string {
if value == nil {
return ""
}
return *value
}
func resolveModelConfig(tenantID string, modelType entity.ModelType, modelRef string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
modelDAO := dao.NewTenantModelDAO()
if _, err := modelDAO.GetByID(modelRef); err == nil {
return resolveModelConfigByID(tenantID, modelType, modelRef)
} else if !errorsIsRecordNotFound(err) {
return nil, "", nil, 0, err
}
return resolveModelConfigFromProviderInstance(tenantID, modelType, modelRef)
}
func resolveModelConfigByID(tenantID string, modelType entity.ModelType, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
modelDAO := dao.NewTenantModelDAO()
instanceDAO := dao.NewTenantModelInstanceDAO()
providerDAO := dao.NewTenantModelProviderDAO()
modelObj, err := modelDAO.GetByID(modelID)
if err != nil {
return nil, "", nil, 0, err
}
if modelObj.Status != "active" {
return nil, "", nil, 0, fmt.Errorf("model %q is disabled", modelID)
}
if !entity.ModelType(modelObj.ModelType).Has(modelType) {
return nil, "", nil, 0, fmt.Errorf("model %q cannot be used as %s model", modelID, modelType.String())
}
instance, err := instanceDAO.GetByID(modelObj.InstanceID)
if err != nil {
return nil, "", nil, 0, err
}
provider, err := providerDAO.GetByID(modelObj.ProviderID)
if err != nil {
return nil, "", nil, 0, err
}
if provider.TenantID != tenantID {
return nil, "", nil, 0, fmt.Errorf("tenant %s has no access to provider owned by tenant %s", tenantID, provider.TenantID)
}
apiKey := instance.APIKey
var extra map[string]string
_ = json.Unmarshal([]byte(instance.Extra), &extra)
region := extra["region"]
baseURL := extra["base_url"]
providerInfo := dao.GetModelProviderManager().FindProvider(provider.ProviderName)
if providerInfo == nil {
return nil, "", nil, 0, fmt.Errorf("provider %q driver not found", provider.ProviderName)
}
driver, err := newModelDriverForBaseURLLocal(providerInfo.ModelDriver, provider.ProviderName, region, baseURL)
if err != nil {
return nil, "", nil, 0, err
}
maxTokens := 0
if mi, _ := dao.GetModelProviderManager().GetModelByName(provider.ProviderName, modelObj.ModelName); mi != nil && mi.MaxTokens != nil {
maxTokens = *mi.MaxTokens
}
if strings.TrimSpace(modelObj.Extra) != "" {
var tenantExtra tenantModelExtra
if err := json.Unmarshal([]byte(modelObj.Extra), &tenantExtra); err != nil {
return nil, "", nil, 0, err
}
if tenantExtra.MaxTokens != nil && *tenantExtra.MaxTokens > 0 {
maxTokens = *tenantExtra.MaxTokens
}
}
apiConfig := &modelModule.APIConfig{ApiKey: &apiKey, Region: &region, BaseURL: &baseURL}
return driver, modelObj.ModelName, apiConfig, maxTokens, nil
}
func resolveModelConfigFromProviderInstance(tenantID string, modelType entity.ModelType, modelName string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
@@ -405,9 +510,9 @@ func parseCompositeModelName(compositeName string) (modelName, instanceName, pro
return parts[0], "default", parts[1], nil
case 1:
return parts[0], "", "", fmt.Errorf("provider name missing in model name: %s", compositeName)
default:
return "", "", "", fmt.Errorf("invalid model name format: %s", compositeName)
}
n := len(parts)
return strings.Join(parts[:n-2], "@"), parts[n-2], parts[n-1], nil
}
func newModelDriverForBaseURLLocal(driver modelModule.ModelDriver, providerName, region, baseURL string) (modelModule.ModelDriver, error) {

View File

@@ -152,16 +152,24 @@ func (s *ChatService) Create(userID string, req map[string]interface{}) (map[str
if llmIDValue, ok := req["llm_id"]; ok {
llmID := stringFromValue(llmIDValue)
llmSetting, _ := mapFromValue(req["llm_setting"])
if err = validateCreateLLMID(llmID, userID, llmSetting); err != nil {
tenantLLMID, err := resolveCreateLLMID(llmID, userID, llmSetting)
if err != nil {
return nil, common.CodeDataError, err
}
if tenantLLMID != "" {
req["tenant_llm_id"] = tenantLLMID
}
}
if rerankIDValue, ok := req["rerank_id"]; ok {
rerankID := stringFromValue(rerankIDValue)
if err = validateCreateRerankID(rerankID, userID); err != nil {
tenantRerankID, err := resolveCreateRerankID(rerankID, userID)
if err != nil {
return nil, common.CodeDataError, err
}
if tenantRerankID != "" {
req["tenant_rerank_id"] = tenantRerankID
}
}
if promptConfigValue, ok := req["prompt_config"]; ok {
@@ -181,6 +189,19 @@ func (s *ChatService) Create(userID string, req map[string]interface{}) (map[str
}
if _, ok := req["llm_id"]; !ok || req["llm_id"] == nil {
req["llm_id"] = tenant.LLMID
if tenant.TenantLLMID != nil {
req["tenant_llm_id"] = *tenant.TenantLLMID
}
}
if stringFromValue(req["llm_id"]) != "" && !isTruthy(req["tenant_llm_id"]) {
llmSetting, _ := mapFromValue(req["llm_setting"])
tenantLLMID, err := resolveCreateLLMID(stringFromValue(req["llm_id"]), userID, llmSetting)
if err != nil {
return nil, common.CodeDataError, err
}
if tenantLLMID != "" {
req["tenant_llm_id"] = tenantLLMID
}
}
if _, ok := req["llm_setting"]; !ok || req["llm_setting"] == nil {
req["llm_setting"] = map[string]interface{}{}
@@ -196,6 +217,7 @@ func (s *ChatService) Create(userID string, req map[string]interface{}) (map[str
}
if _, ok := req["rerank_id"]; !ok {
req["rerank_id"] = ""
req["tenant_rerank_id"] = nil
}
if _, ok := req["similarity_threshold"]; !ok {
req["similarity_threshold"] = 0.1
@@ -295,9 +317,9 @@ func (s *ChatService) validateCreateDatasetIDs(value interface{}, tenantID strin
return normalizedIDs, nil
}
func validateCreateLLMID(llmID, tenantID string, llmSetting map[string]interface{}) error {
func resolveCreateLLMID(llmID, tenantID string, llmSetting map[string]interface{}) (string, error) {
if llmID == "" {
return nil
return "", nil
}
modelType := entity.ModelTypeChat
switch confModelType := llmSetting["model_type"].(type) {
@@ -320,24 +342,34 @@ func validateCreateLLMID(llmID, tenantID string, llmSetting map[string]interface
}
}
}
if _, _, _, _, err := NewModelProviderService().GetModelConfigFromProviderInstance(tenantID, modelType, llmID); err != nil {
return fmt.Errorf("`llm_id` %s doesn't exist", llmID)
modelProvider := NewModelProviderService()
if _, _, _, _, err := modelProvider.ResolveModelConfig(tenantID, modelType, llmID); err != nil {
return "", fmt.Errorf("`llm_id` %s doesn't exist", llmID)
}
return nil
tenantLLMID, err := modelProvider.ResolveModelID(tenantID, modelType, llmID)
if err != nil {
return "", err
}
return tenantLLMID, nil
}
func validateCreateRerankID(rerankID, tenantID string) error {
func resolveCreateRerankID(rerankID, tenantID string) (string, error) {
if rerankID == "" {
return nil
return "", nil
}
llmName := strings.Split(rerankID, "@")[0]
if _, ok := DefaultRerankModels[llmName]; ok {
return nil
return "", nil
}
if _, _, _, _, err := NewModelProviderService().GetModelConfigFromProviderInstance(tenantID, entity.ModelTypeRerank, rerankID); err != nil {
return fmt.Errorf("`rerank_id` %s doesn't exist", rerankID)
modelProvider := NewModelProviderService()
if _, _, _, _, err := modelProvider.ResolveModelConfig(tenantID, entity.ModelTypeRerank, rerankID); err != nil {
return "", fmt.Errorf("`rerank_id` %s doesn't exist", rerankID)
}
return nil
tenantRerankID, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeRerank, rerankID)
if err != nil {
return "", err
}
return tenantRerankID, nil
}
func applyCreatePromptDefaults(req map[string]interface{}) {
@@ -398,6 +430,8 @@ func buildCreateChatEntity(req map[string]interface{}, tenantID string) *entity.
icon := stringFromValue(req["icon"])
llmID := stringFromValue(req["llm_id"])
rerankID := stringFromValue(req["rerank_id"])
tenantLLMID := stringFromValue(req["tenant_llm_id"])
tenantRerankID := stringFromValue(req["tenant_rerank_id"])
llmSetting, _ := mapFromValue(req["llm_setting"])
promptConfig, _ := mapFromValue(req["prompt_config"])
kbIDs, _ := stringListFromValue(req["kb_ids"])
@@ -418,6 +452,7 @@ func buildCreateChatEntity(req map[string]interface{}, tenantID string) *entity.
Description: &description,
Icon: &icon,
LLMID: llmID,
TenantLLMID: stringPtrIfNotEmpty(tenantLLMID),
LLMSetting: entity.JSONMap(llmSetting),
PromptType: stringFromValue(req["prompt_type"]),
PromptConfig: entity.JSONMap(promptConfig),
@@ -427,6 +462,7 @@ func buildCreateChatEntity(req map[string]interface{}, tenantID string) *entity.
TopK: int64FromValue(req["top_k"]),
DoRefer: stringFromValue(req["do_refer"]),
RerankID: rerankID,
TenantRerankID: stringPtrIfNotEmpty(tenantRerankID),
KBIDs: kbIDsJSON,
Status: &statusValue,
}
@@ -449,6 +485,13 @@ func buildCreateChatEntity(req map[string]interface{}, tenantID string) *entity.
return chat
}
func stringPtrIfNotEmpty(value string) *string {
if value == "" {
return nil
}
return &value
}
func (s *ChatService) buildCreateChatResponse(chat *entity.Chat) (map[string]interface{}, error) {
data, err := structToMap(chat)
if err != nil {
@@ -761,16 +804,24 @@ func (s *ChatService) updateChatREST(userID, chatID string, req map[string]inter
if value, ok := req["llm_id"]; ok {
llmID := fmt.Sprint(value)
if err := s.validateRESTLLMID(llmID, userID, llmSetting); err != nil {
tenantLLMID, err := s.resolveRESTLLMID(llmID, userID, llmSetting)
if err != nil {
return nil, err
}
if tenantLLMID != "" {
req["tenant_llm_id"] = tenantLLMID
}
}
if value, ok := req["rerank_id"]; ok {
rerankID := fmt.Sprint(value)
if err := s.validateRESTRerankID(rerankID, userID); err != nil {
tenantRerankID, err := s.resolveRESTRerankID(rerankID, userID)
if err != nil {
return nil, err
}
if tenantRerankID != "" {
req["tenant_rerank_id"] = tenantRerankID
}
}
if value, ok := req["prompt_config"]; ok {
@@ -909,9 +960,9 @@ func (s *ChatService) validateRESTDatasetIDs(value interface{}, userID string) (
return kbIDs, nil
}
func (s *ChatService) validateRESTLLMID(llmID, tenantID string, llmSetting map[string]interface{}) error {
func (s *ChatService) resolveRESTLLMID(llmID, tenantID string, llmSetting map[string]interface{}) (string, error) {
if llmID == "" {
return nil
return "", nil
}
modelType := entity.ModelTypeChat
if rawModelType, ok := llmSetting["model_type"]; ok {
@@ -929,24 +980,34 @@ func (s *ChatService) validateRESTLLMID(llmID, tenantID string, llmSetting map[s
}
}
}
if _, _, _, _, err := NewModelProviderService().GetModelConfigFromProviderInstance(tenantID, modelType, llmID); err != nil {
return fmt.Errorf("`llm_id` %s doesn't exist", llmID)
modelProvider := NewModelProviderService()
if _, _, _, _, err := modelProvider.ResolveModelConfig(tenantID, modelType, llmID); err != nil {
return "", fmt.Errorf("`llm_id` %s doesn't exist", llmID)
}
return nil
tenantLLMID, err := modelProvider.ResolveModelID(tenantID, modelType, llmID)
if err != nil {
return "", err
}
return tenantLLMID, nil
}
func (s *ChatService) validateRESTRerankID(rerankID, tenantID string) error {
func (s *ChatService) resolveRESTRerankID(rerankID, tenantID string) (string, error) {
if rerankID == "" {
return nil
return "", nil
}
baseName := s.splitModelNameAndFactory(rerankID)
if _, ok := defaultRerankModels[baseName]; ok {
return nil
return "", nil
}
if _, _, _, _, err := NewModelProviderService().GetModelConfigFromProviderInstance(tenantID, entity.ModelTypeRerank, rerankID); err != nil {
return fmt.Errorf("`rerank_id` %s doesn't exist", rerankID)
modelProvider := NewModelProviderService()
if _, _, _, _, err := modelProvider.ResolveModelConfig(tenantID, entity.ModelTypeRerank, rerankID); err != nil {
return "", fmt.Errorf("`rerank_id` %s doesn't exist", rerankID)
}
return nil
tenantRerankID, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeRerank, rerankID)
if err != nil {
return "", err
}
return tenantRerankID, nil
}
func filterRESTChatUpdates(req map[string]interface{}) map[string]interface{} {

View File

@@ -1866,7 +1866,7 @@ func (s *ChatPipelineService) getLLMModelConfig(chat *entity.Chat) (map[string]i
// when the LLM is registered as such, otherwise CHAT.
modelType := entity.ModelTypeChat
modelTypeStr := "chat"
if modelTypes, mtErr := s.ModelProviderSvc.GetModelTypeByName(chat.TenantID, chat.LLMID); mtErr == nil {
if modelTypes, mtErr := s.ModelProviderSvc.ResolveModelType(chat.TenantID, chat.LLMID); mtErr == nil {
for _, mt := range modelTypes {
if mt == entity.ModelTypeImage2Text {
modelType = entity.ModelTypeImage2Text
@@ -1876,7 +1876,7 @@ func (s *ChatPipelineService) getLLMModelConfig(chat *entity.Chat) (map[string]i
}
}
cfg, modelName, factoryName, baseURL, err := s.buildLLMModelConfig(
s.ModelProviderSvc.GetModelConfigFromProviderInstance(chat.TenantID, modelType, chat.LLMID),
s.ModelProviderSvc.ResolveModelConfig(chat.TenantID, modelType, chat.LLMID),
)
if err != nil {
return nil, "", "", "", err
@@ -1955,7 +1955,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat)
}
if kbs[0].EmbdID != "" {
embdTenantID := kbs[0].TenantID
driver, modelName, apiConfig, maxTokens, err := s.ModelProviderSvc.GetModelConfigFromProviderInstance(
driver, modelName, apiConfig, maxTokens, err := s.ModelProviderSvc.ResolveModelConfig(
embdTenantID, entity.ModelTypeEmbedding, kbs[0].EmbdID,
)
if err != nil {
@@ -1979,7 +1979,7 @@ func (s *ChatPipelineService) getModels(ctx context.Context, chat *entity.Chat)
// Rerank model.
var rerankModel *modelModule.RerankModel
if chat.RerankID != "" {
rerankDriver, rerankName, rerankConfig, _, err := s.ModelProviderSvc.GetModelConfigFromProviderInstance(
rerankDriver, rerankName, rerankConfig, _, err := s.ModelProviderSvc.ResolveModelConfig(
chat.TenantID, entity.ModelTypeRerank, chat.RerankID,
)
if err == nil {

View File

@@ -21,7 +21,15 @@ func setupChatRESTUpdateServiceTestDB(t *testing.T) *gorm.DB {
t.Fatalf("failed to open sqlite: %v", err)
}
if err := db.AutoMigrate(&entity.Chat{}, &entity.Tenant{}, &entity.Knowledgebase{}, &entity.UserTenant{}); err != nil {
if err := db.AutoMigrate(
&entity.Chat{},
&entity.Tenant{},
&entity.Knowledgebase{},
&entity.UserTenant{},
&entity.TenantModelProvider{},
&entity.TenantModelInstance{},
&entity.TenantModel{},
); err != nil {
t.Fatalf("failed to migrate test schema: %v", err)
}
@@ -42,6 +50,34 @@ func setupChatRESTUpdateServiceTestDB(t *testing.T) *gorm.DB {
}).Error; err != nil {
t.Fatalf("failed to create tenant: %v", err)
}
if err := db.Create(&entity.TenantModelProvider{
ID: "provider-a",
TenantID: "user-1",
ProviderName: "OpenAI",
}).Error; err != nil {
t.Fatalf("failed to create model provider: %v", err)
}
if err := db.Create(&entity.TenantModelInstance{
ID: "instance-a",
ProviderID: "provider-a",
InstanceName: "default",
APIKey: "sk-test",
Status: "active",
Extra: "{}",
}).Error; err != nil {
t.Fatalf("failed to create model instance: %v", err)
}
if err := db.Create(&entity.TenantModel{
ID: "model-a",
ProviderID: "provider-a",
InstanceID: "instance-a",
ModelName: "gpt-test",
ModelType: int(entity.ModelTypeChat),
Status: "active",
Extra: "{}",
}).Error; err != nil {
t.Fatalf("failed to create tenant model: %v", err)
}
return db
}

View File

@@ -247,7 +247,7 @@ func (s *ChunkService) RetrievalTest(req *service.RetrievalTestRequest, userID s
modelProviderSvc := service.NewModelProviderService()
if chatID != "" {
// Use chat_id from search_config (it's actually the model name)
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeChat, chatID)
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, chatID)
if getErr != nil {
common.Warn("Failed to get chat model from search_config chat_id, using tenant default", zap.String("chatID", chatID), zap.Error(getErr))
} else {
@@ -266,7 +266,7 @@ func (s *ChunkService) RetrievalTest(req *service.RetrievalTestRequest, userID s
if err != nil || modelName == "" {
common.Warn("Failed to get tenant default chat model name for meta_data_filter", zap.Error(err))
} else {
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeChat, modelName)
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, modelName)
if getErr != nil {
common.Warn("Failed to get chat model for meta_data_filter", zap.Error(getErr))
} else {
@@ -312,7 +312,7 @@ func (s *ChunkService) RetrievalTest(req *service.RetrievalTestRequest, userID s
if err != nil || llmModelName == "" {
common.Warn("Failed to get default chat model name for LLM transformations", zap.Error(err))
} else {
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeChat, llmModelName)
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, llmModelName)
if getErr != nil {
common.Warn("Failed to get chat model for LLM transformations", zap.Error(getErr))
} else {
@@ -362,20 +362,20 @@ func (s *ChunkService) RetrievalTest(req *service.RetrievalTestRequest, userID s
var embeddingModel *models.EmbeddingModel
var embdID string
if kbRecords[0].TenantEmbdID != nil && *kbRecords[0].TenantEmbdID != "" {
driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.GetModelConfigByID(tenantIDs[0], *kbRecords[0].TenantEmbdID)
driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.GetModelConfigByID(tenantIDs[0], entity.ModelTypeEmbedding, *kbRecords[0].TenantEmbdID)
if getErr != nil {
return nil, fmt.Errorf("failed to get embedding model by tenant_embd_id: %w", getErr)
}
embeddingModel = models.NewEmbeddingModel(driver, &modelName, apiConfig, maxTokens)
} else if kbRecords[0].EmbdID != "" {
embdID = kbRecords[0].EmbdID
driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeEmbedding, embdID)
driver, modelName, apiConfig, maxTokens, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeEmbedding, embdID)
if getErr != nil {
_, embdID, err = dao.LookupTenantLLMByName(dao.NewTenantLLMDAO(), tenantIDs[0], kbRecords[0].EmbdID, entity.ModelTypeEmbedding)
if err != nil {
return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", getErr)
}
driver, modelName, apiConfig, maxTokens, getErr = modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeEmbedding, embdID)
driver, modelName, apiConfig, maxTokens, getErr = modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeEmbedding, embdID)
if getErr != nil {
return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", getErr)
}
@@ -401,14 +401,14 @@ func (s *ChunkService) RetrievalTest(req *service.RetrievalTestRequest, userID s
// Get rerank model if RerankID is specified
var rerankModel *models.RerankModel
if req.TenantRerankID != nil && *req.TenantRerankID != "" {
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.GetModelConfigByID(tenantIDs[0], *req.TenantRerankID)
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.GetModelConfigByID(tenantIDs[0], entity.ModelTypeRerank, *req.TenantRerankID)
if getErr != nil {
return nil, fmt.Errorf("failed to get rerank model by tenant_rerank_id: %w", getErr)
}
rerankModel = models.NewRerankModel(driver, &mdlName, apiConfig)
} else if req.RerankID != nil && *req.RerankID != "" {
rerankCompositeName := *req.RerankID
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeRerank, rerankCompositeName)
driver, mdlName, apiConfig, _, getErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeRerank, rerankCompositeName)
if getErr != nil {
rerankModel = nil
} else {

View File

@@ -905,7 +905,7 @@ func (d *DatasetService) CheckEmbedding(userID, datasetID string, req *CheckEmbe
return nil, common.CodeServerError, errors.New("doc engine not initialized")
}
driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().GetModelConfigFromProviderInstance(kb.TenantID, entity.ModelTypeEmbedding, embeddingID)
driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().ResolveModelConfig(kb.TenantID, entity.ModelTypeEmbedding, embeddingID)
if err != nil {
return nil, common.CodeDataError, err
}
@@ -1990,7 +1990,7 @@ func (d *DatasetService) SearchDatasets(req *SearchDatasetsRequest, userID strin
method, _ := metadataFilter["method"].(string)
if method == "auto" || method == "semi_auto" {
if chatID != "" {
driver, modelName, apiConfig, _, err := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeChat, chatID)
driver, modelName, apiConfig, _, err := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeChat, chatID)
if err != nil {
common.Warn("Failed to get chat model config from search_config chat_id, using tenant default", zap.String("chatID", chatID), zap.Error(err))
} else {
@@ -2081,7 +2081,7 @@ func (d *DatasetService) SearchDatasets(req *SearchDatasetsRequest, userID strin
// Determine embedding model
var embeddingModel *models.EmbeddingModel
if kbRecords[0].EmbdID != "" {
driver, modelName, apiConfig, maxTokens, embErr := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeEmbedding, kbRecords[0].EmbdID)
driver, modelName, apiConfig, maxTokens, embErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeEmbedding, kbRecords[0].EmbdID)
if embErr != nil {
return nil, fmt.Errorf("failed to get embedding model by embd_id: %w", embErr)
}
@@ -2095,7 +2095,7 @@ func (d *DatasetService) SearchDatasets(req *SearchDatasetsRequest, userID strin
// Get rerank model if rerankID is specified
var rerankModel *models.RerankModel
if rerankID != "" {
driver, modelName, apiConfig, _, rErr := modelProviderSvc.GetModelConfigFromProviderInstance(tenantIDs[0], entity.ModelTypeRerank, rerankID)
driver, modelName, apiConfig, _, rErr := modelProviderSvc.ResolveModelConfig(tenantIDs[0], entity.ModelTypeRerank, rerankID)
if rErr != nil {
return nil, fmt.Errorf("failed to get rerank model by rerank_id: %w", rErr)
}
@@ -2469,12 +2469,22 @@ func (d *DatasetService) CreateDataset(req *CreateDatasetRequest, tenantID strin
parserConfig["llm_id"] = tenant.LLMID
embdID := tenant.EmbdID
tenantEmbdID := ptrStringValue(tenant.TenantEmbdID)
if embeddingModel != "" {
ok, message := d.verifyEmbeddingAvailability(embeddingModel, tenantID)
if !ok {
return nil, common.CodeDataError, errors.New(message)
}
embdID = embeddingModel
tenantEmbdID = ""
}
if embdID != "" && tenantEmbdID == "" {
resolvedID, err := NewModelProviderService().ResolveModelID(tenantID, entity.ModelTypeEmbedding, embdID)
if err == nil {
tenantEmbdID = resolvedID
} else {
return nil, common.CodeDataError, err
}
}
kbID := utility.GenerateToken()
@@ -2499,6 +2509,7 @@ func (d *DatasetService) CreateDataset(req *CreateDatasetRequest, tenantID strin
ParserConfig: parserConfig,
Permission: permission,
EmbdID: embdID,
TenantEmbdID: stringPtrIfNotEmpty(tenantEmbdID),
Status: &status,
}
@@ -2785,14 +2796,24 @@ func (d *DatasetService) UpdateDataset(datasetID, tenantID string, req UpdateDat
}
}
if embdIDProvided {
tenantEmbdID := ptrStringValue(kb.TenantEmbdID)
if embdID == "" {
embdID = kb.EmbdID
} else {
tenantEmbdID = ""
}
ok, message := d.verifyEmbeddingAvailability(embdID, tenantID)
if !ok {
return nil, common.CodeDataError, errors.New(message)
}
if embdID != "" && tenantEmbdID == "" {
resolvedID, err := NewModelProviderService().ResolveModelID(tenantID, entity.ModelTypeEmbedding, embdID)
if err == nil {
tenantEmbdID = resolvedID
}
}
updates["embd_id"] = embdID
updates["tenant_embd_id"] = stringPtrIfNotEmpty(tenantEmbdID)
}
if req.AutoMetadataConfig != nil {
@@ -3572,14 +3593,23 @@ func validateDatasetAvatar(avatar string) error {
func validateDatasetEmbeddingModel(embeddingModel string) error {
if embeddingModel == "" {
return errors.New("Embedding model identifier must follow <model_name>@<provider> format")
return errors.New("Embedding model identifier is required")
}
modelName, provider, ok := strings.Cut(embeddingModel, "@")
if !ok {
if !strings.Contains(embeddingModel, "@") {
return nil
}
parts := strings.Split(embeddingModel, "@")
for _, part := range parts {
if strings.TrimSpace(part) == "" {
return errors.New("Both model_name and provider must be non-empty strings")
}
}
if len(parts) < 2 {
return errors.New("Embedding model identifier must follow <model_name>@<provider> format")
}
if strings.TrimSpace(modelName) == "" || strings.TrimSpace(provider) == "" {
if strings.TrimSpace(parts[0]) == "" || strings.TrimSpace(parts[len(parts)-1]) == "" {
return errors.New("Both model_name and provider must be non-empty strings")
}
@@ -3637,7 +3667,7 @@ func normalizeDatasetID(id string) (string, error) {
}
func (d *DatasetService) verifyEmbeddingAvailability(embdID string, tenantID string) (bool, string) {
_, _, _, _, err := NewModelProviderService().GetModelConfigFromProviderInstance(tenantID, entity.ModelTypeEmbedding, embdID)
_, _, _, _, err := NewModelProviderService().ResolveModelConfig(tenantID, entity.ModelTypeEmbedding, embdID)
if err != nil {
return false, err.Error()
}

View File

@@ -296,6 +296,40 @@ func TestDatasetServiceUpdateDatasetAcceptsProviderInstanceEmbedding(t *testing.
}
}
func TestDatasetServiceUpdateDatasetAcceptsEmbeddingModelID(t *testing.T) {
db := setupDatasetUpdateTestDB(t)
pushServiceDB(t, db)
insertDatasetUpdateKB(t, "kb-1", "tenant-1", "Original")
insertDatasetUpdateModelProvider(t, "provider-1", "tenant-1", "ZHIPU-AI")
insertDatasetUpdateModelInstance(t, "instance-1", "provider-1", "test")
insertDatasetUpdateTenantModel(t, "model-1", "provider-1", "instance-1", "embedding-2", int(entity.ModelTypeEmbedding))
embeddingModelID := "model-1"
result, code, err := testDatasetUpdateService(t).UpdateDataset("kb-1", "tenant-1", UpdateDatasetRequest{
EmbeddingModel: &embeddingModelID,
})
if err != nil {
t.Fatalf("UpdateDataset failed: %v", err)
}
if code != common.CodeSuccess {
t.Fatalf("expected success code, got %d", code)
}
if result["embedding_model"] != embeddingModelID {
t.Fatalf("expected embedding model %q, got %#v", embeddingModelID, result["embedding_model"])
}
persisted, err := dao.NewKnowledgebaseDAO().GetByID("kb-1")
if err != nil {
t.Fatalf("get updated kb: %v", err)
}
if persisted.EmbdID != embeddingModelID {
t.Fatalf("expected persisted embedding model %q, got %q", embeddingModelID, persisted.EmbdID)
}
if persisted.TenantEmbdID == nil || *persisted.TenantEmbdID != embeddingModelID {
t.Fatalf("expected persisted tenant_embd_id %q, got %#v", embeddingModelID, persisted.TenantEmbdID)
}
}
func TestDatasetServiceUpdateDatasetRejectsEmptyConnectorID(t *testing.T) {
db := setupDatasetUpdateTestDB(t)
pushServiceDB(t, db)

View File

@@ -111,7 +111,7 @@ func CrossLanguages(ctx context.Context, tenantID string, llmID string, query st
var err error
if llmID != "" {
modelTypes, err := modelProviderSvc.GetModelTypeByName(tenantID, llmID)
modelTypes, err := modelProviderSvc.ResolveModelType(tenantID, llmID)
if err != nil {
return query, fmt.Errorf("failed to get model type: %w", err)
}
@@ -122,7 +122,7 @@ func CrossLanguages(ctx context.Context, tenantID string, llmID string, query st
break
}
}
driver, modelName, apiConfig, _, err := modelProviderSvc.GetModelConfigFromProviderInstance(tenantID, resolvedType, llmID)
driver, modelName, apiConfig, _, err := modelProviderSvc.ResolveModelConfig(tenantID, resolvedType, llmID)
if err != nil {
return query, fmt.Errorf("failed to get chat model: %w", err)
}

View File

@@ -343,23 +343,20 @@ type ListMemoryResponse struct {
// req := &CreateMemoryRequest{Name: "MyMemory", MemoryType: []string{"semantic"}, EmbdID: "embd1", LLMID: "llm1"}
// resp, err := service.CreateMemory("tenant123", req)
func (s *MemoryService) CreateMemory(tenantID string, req *CreateMemoryRequest) (*CreateMemoryResponse, error) {
// Ensure tenant model IDs are populated for LLM and embedding model parameters
// This automatically fills tenant_llm_id and tenant_embd_id based on llm_id and embd_id
tenantLLMService := NewTenantLLMService()
params := map[string]interface{}{
"llm_id": req.LLMID,
"embd_id": req.EmbdID,
modelProvider := NewModelProviderService()
if req.LLMID != "" && req.TenantLLMID == nil {
tenantLLMID, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeChat, req.LLMID)
if err != nil {
return nil, err
}
req.TenantLLMID = &tenantLLMID
}
params = tenantLLMService.EnsureTenantModelIDForParams(tenantID, params)
// Update request with tenant model IDs from the processed params
if tenantLLMID, ok := params["tenant_llm_id"].(int64); ok {
tenantLLMIDStr := strconv.FormatInt(tenantLLMID, 10)
req.TenantLLMID = &tenantLLMIDStr
}
if tenantEmbdID, ok := params["tenant_embd_id"].(int64); ok {
tenantEmbdIDStr := strconv.FormatInt(tenantEmbdID, 10)
req.TenantEmbdID = &tenantEmbdIDStr
if req.EmbdID != "" && req.TenantEmbdID == nil {
tenantEmbdID, err := modelProvider.ResolveModelID(tenantID, entity.ModelTypeEmbedding, req.EmbdID)
if err != nil {
return nil, err
}
req.TenantEmbdID = &tenantEmbdID
}
if err := common.ValidateName(req.Name); err != nil {
@@ -1261,7 +1258,7 @@ func memoryMessageTextExpr(question string, similarityThreshold float64) *engine
}
func (s *MemoryService) memoryMessageDenseExpr(question string, memory *entity.Memory, topN int, similarityThreshold float64) (*enginetypes.MatchDenseExpr, error) {
driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().GetModelConfigFromProviderInstance(memory.TenantID, entity.ModelTypeEmbedding, memory.EmbdID)
driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().ResolveModelConfig(memory.TenantID, entity.ModelTypeEmbedding, memory.EmbdID)
if err != nil {
return nil, err
}

View File

@@ -262,7 +262,7 @@ func (s *MemoryMessageService) embedAndSave(ctx context.Context, mem *CreateMemo
}
content, _ := rawMessage["content"].(string)
driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().GetModelConfigFromProviderInstance(mem.TenantID, entity.ModelTypeEmbedding, mem.EmbdID)
driver, modelName, apiConfig, maxTokens, err := NewModelProviderService().ResolveModelConfig(mem.TenantID, entity.ModelTypeEmbedding, mem.EmbdID)
if err != nil {
return err
}

View File

@@ -130,6 +130,7 @@ func NewModelProviderService() *ModelProviderService {
modelDAO: dao.NewTenantModelDAO(),
modelGroupDAO: dao.NewTenantModelGroupDAO(),
modelGroupMappingDAO: dao.NewTenantModelGroupMappingDAO(),
tenantDAO: dao.NewTenantDAO(),
userTenantDAO: dao.NewUserTenantDAO(),
}
}
@@ -140,6 +141,7 @@ type ModelProviderService struct {
modelDAO *dao.TenantModelDAO
modelGroupDAO *dao.TenantModelGroupDAO
modelGroupMappingDAO *dao.TenantModelGroupMappingDAO
tenantDAO *dao.TenantDAO
userTenantDAO *dao.UserTenantDAO
}
@@ -811,20 +813,19 @@ func (m *ModelProviderService) ShowTask(providerName, instanceName, taskID, user
// to ListTenantDefaultModels (which only enumerates the 6-7 default
// tenant fields and returned `[]` for any tenant without defaults),
// breaking the front-end's "View Models" list entirely.
func (m *ModelProviderService) ListTenantAddedModels(userID, modelTypeFilter string) ([]map[string]interface{}, common.ErrorCode, error) {
// Resolve tenant. Match the convention used elsewhere in this file
// (see ListProviderInstances, DropProviderInstances): take the first
// tenant where the user has role=owner.
tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner")
func (m *ModelProviderService) ListTenantAddedModels(userID, ownerTenantID, modelTypeFilter string) ([]map[string]interface{}, common.ErrorCode, error) {
tenant, code, err := m.resolveModelListTenant(userID, ownerTenantID)
if err != nil {
return nil, common.CodeServerError, err
return nil, code, err
}
if len(tenants) == 0 {
// No tenant for the user → empty list, code=0. Python returns
// get_result(data=[]) for the same path.
if tenant == nil {
return []map[string]interface{}{}, common.CodeSuccess, nil
}
tenantID := tenants[0].TenantID
tenantID := tenant.ID
tenantName := ""
if tenant.Name != nil {
tenantName = *tenant.Name
}
if modelTypeFilter != "" {
modelTypeFilter = strings.ToLower(strings.TrimSpace(modelTypeFilter))
@@ -869,12 +870,16 @@ func (m *ModelProviderService) ListTenantAddedModels(userID, modelTypeFilter str
}
activeByKey := make(map[string]int)
inactiveByKey := make(map[string]int)
activeModelIDByKey := make(map[string]string)
for _, rec := range modelRecords {
key := rec.ProviderID + "@" + rec.InstanceID + "@" + rec.ModelName
if rec.Status == "inactive" {
inactiveByKey[key] |= rec.ModelType
} else {
activeByKey[key] |= rec.ModelType
if activeModelIDByKey[key] == "" {
activeModelIDByKey[key] = rec.ID
}
}
}
@@ -930,7 +935,10 @@ func (m *ModelProviderService) ListTenantAddedModels(userID, modelTypeFilter str
continue
}
added = append(added, map[string]interface{}{
"model_type": merged.HumanReadable(),
"model_id": activeModelIDByKey[key],
"tenant_id": tenant.ID,
"tenant_name": tenantName,
"model_type": modelTypesForAPI(merged),
"name": llm.Name,
"provider_id": inst.ProviderID,
"provider_name": p.ProviderName,
@@ -944,6 +952,53 @@ func (m *ModelProviderService) ListTenantAddedModels(userID, modelTypeFilter str
return added, common.CodeSuccess, nil
}
func modelTypesForAPI(modelType entity.ModelType) []string {
types := modelType.HumanReadable()
for i, t := range types {
if t == "image2text" {
types[i] = "vision"
}
}
return types
}
func (m *ModelProviderService) resolveModelListTenant(userID, ownerTenantID string) (*entity.Tenant, common.ErrorCode, error) {
if ownerTenantID == "" {
tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner")
if err != nil {
return nil, common.CodeServerError, err
}
if len(tenants) == 0 {
return nil, common.CodeSuccess, nil
}
ownerTenantID = tenants[0].TenantID
} else {
relations, err := m.userTenantDAO.GetByUserID(userID)
if err != nil {
return nil, common.CodeServerError, err
}
allowed := false
for _, rel := range relations {
if rel.TenantID == ownerTenantID {
allowed = true
break
}
}
if !allowed {
return nil, common.CodeAuthenticationError, fmt.Errorf("permission denied")
}
}
tenant, err := m.tenantDAO.GetByID(ownerTenantID)
if err != nil {
if dao.IsNotFoundErr(err) {
return nil, common.CodeNotFound, fmt.Errorf("tenant %s not found", ownerTenantID)
}
return nil, common.CodeServerError, err
}
return tenant, common.CodeSuccess, nil
}
func (m *ModelProviderService) AlterProviderInstance(userID, providerName, instanceName, newInstanceName, apiKey string) (common.ErrorCode, error) {
return common.CodeSuccess, nil
}
@@ -2087,7 +2142,7 @@ func (m *ModelProviderService) ParseFile(providerName, instanceName, modelName,
// GetEmbeddingModel returns an EmbeddingModel wrapper for the given tenant
func (m *ModelProviderService) GetEmbeddingModel(tenantID, compositeModelName string) (*modelModule.EmbeddingModel, error) {
driver, modelName, apiConfig, maxTokens, err := m.getModelConfig(tenantID, compositeModelName)
driver, modelName, apiConfig, maxTokens, err := m.ResolveModelConfig(tenantID, entity.ModelTypeEmbedding, compositeModelName)
if err != nil {
return nil, err
}
@@ -2096,7 +2151,7 @@ func (m *ModelProviderService) GetEmbeddingModel(tenantID, compositeModelName st
// GetChatModel returns a ChatModel wrapper for the given tenant
func (m *ModelProviderService) GetChatModel(tenantID, compositeModelName string) (*modelModule.ChatModel, error) {
driver, modelName, apiConfig, _, err := m.getModelConfig(tenantID, compositeModelName)
driver, modelName, apiConfig, _, err := m.ResolveModelConfig(tenantID, entity.ModelTypeChat, compositeModelName)
if err != nil {
return nil, err
}
@@ -2105,7 +2160,7 @@ func (m *ModelProviderService) GetChatModel(tenantID, compositeModelName string)
// GetRerankModel returns a RerankModel wrapper for the given tenant
func (m *ModelProviderService) GetRerankModel(tenantID, compositeModelName string) (*modelModule.RerankModel, error) {
driver, modelName, apiConfig, _, err := m.getModelConfig(tenantID, compositeModelName)
driver, modelName, apiConfig, _, err := m.ResolveModelConfig(tenantID, entity.ModelTypeRerank, compositeModelName)
if err != nil {
return nil, err
}
@@ -2123,22 +2178,33 @@ func (m *ModelProviderService) GetTenantDefaultModelByType(tenantID string, mode
return nil, "", nil, 0, fmt.Errorf("OCR model name is required")
}
tenantSvc := NewTenantService()
modelName, err := tenantSvc.GetDefaultModelName(tenantID, modelType)
tenant, err := m.tenantDAO.GetByID(tenantID)
if err != nil {
return nil, "", nil, 0, fmt.Errorf("failed to get default model name for tenant: %s type %s: %w", tenantID, modelType, err)
return nil, "", nil, 0, fmt.Errorf("failed to get tenant: %s type %s: %w", tenantID, modelType, err)
}
modelName, modelID := defaultModelRefs(tenant, modelType)
if modelID != "" {
driver, resolvedName, apiConfig, maxTokens, idErr := m.GetModelConfigByID(tenantID, modelType, modelID)
if idErr == nil {
return driver, resolvedName, apiConfig, maxTokens, nil
}
common.Warn("GetTenantDefaultModelByType: model_id lookup failed, falling back to model name",
zap.String("tenantID", tenantID),
zap.String("modelID", modelID),
zap.Error(idErr))
}
if modelName == "" {
return nil, "", nil, 0, fmt.Errorf("no default %s model is set", modelType)
}
return m.GetModelConfigFromProviderInstance(tenantID, modelType, modelName)
return m.ResolveModelConfig(tenantID, modelType, modelName)
}
// GetModelConfigByID returns model driver and API config for a tenant_model row by its ID.
func (m *ModelProviderService) GetModelConfigByID(userID, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
func (m *ModelProviderService) GetModelConfigByID(userID string, modelType entity.ModelType, modelID string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
common.Debug("GetModelConfigByID",
zap.String("userID", userID),
zap.String("modelType", modelType.String()),
zap.String("modelID", modelID))
modelEntity, err := m.modelDAO.GetByID(modelID)
@@ -2151,6 +2217,9 @@ func (m *ModelProviderService) GetModelConfigByID(userID, modelID string) (model
if modelEntity.Status != "active" {
return nil, "", nil, 0, fmt.Errorf("tenant model id=%s is disabled", modelID)
}
if modelType != 0 && !entity.ModelType(modelEntity.ModelType).Has(modelType) {
return nil, "", nil, 0, fmt.Errorf("tenant model id=%s cannot be used as %s model", modelID, modelType.String())
}
providerEntity, err := m.modelProviderDAO.GetByID(modelEntity.ProviderID)
if err != nil {
@@ -2220,6 +2289,106 @@ func (m *ModelProviderService) GetModelConfigByID(userID, modelID string) (model
return modelDriver, modelEntity.ModelName, apiConfig, maxTokens, nil
}
func defaultModelRefs(tenant *entity.Tenant, modelType entity.ModelType) (string, string) {
if tenant == nil {
return "", ""
}
switch modelType {
case entity.ModelTypeChat:
return tenant.LLMID, ptrStringValue(tenant.TenantLLMID)
case entity.ModelTypeEmbedding:
return tenant.EmbdID, ptrStringValue(tenant.TenantEmbdID)
case entity.ModelTypeRerank:
return tenant.RerankID, ptrStringValue(tenant.TenantRerankID)
case entity.ModelTypeSpeech2Text:
return tenant.ASRID, ptrStringValue(tenant.TenantASRID)
case entity.ModelTypeImage2Text:
return tenant.Img2TxtID, ptrStringValue(tenant.TenantImg2TxtID)
case entity.ModelTypeTTS:
return tenant.TTSID, ptrStringValue(tenant.TenantTTSID)
case entity.ModelTypeOCR:
return tenant.OCRID, ptrStringValue(tenant.TenantOCRID)
default:
return "", ""
}
}
func (m *ModelProviderService) ResolveModelConfig(tenantID string, modelType entity.ModelType, modelRef string) (modelModule.ModelDriver, string, *modelModule.APIConfig, int, error) {
if strings.TrimSpace(modelRef) == "" {
return nil, "", nil, 0, fmt.Errorf("model ref is required")
}
if _, err := m.modelDAO.GetByID(modelRef); err == nil {
return m.GetModelConfigByID(tenantID, modelType, modelRef)
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, "", nil, 0, err
}
return m.GetModelConfigFromProviderInstance(tenantID, modelType, modelRef)
}
func (m *ModelProviderService) ResolveModelID(tenantID string, modelType entity.ModelType, modelName string) (string, error) {
if modelObj, err := m.modelDAO.GetByID(modelName); err == nil {
if modelObj.Status != "active" {
return "", fmt.Errorf("tenant model id=%s is disabled", modelName)
}
if !entity.ModelType(modelObj.ModelType).Has(modelType) {
return "", fmt.Errorf("tenant model id=%s cannot be used as %s model", modelName, modelType.String())
}
if _, _, _, _, err := m.GetModelConfigByID(tenantID, modelType, modelName); err != nil {
return "", err
}
return modelObj.ID, nil
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return "", err
}
pureModelName, instanceName, providerName, err := parseModelName(modelName)
if err != nil {
return "", err
}
provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName)
if err != nil {
return "", fmt.Errorf("provider %q lookup failed: %w", providerName, err)
}
if provider == nil {
return "", fmt.Errorf("provider %q not found for model %q", providerName, modelName)
}
instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName)
if err != nil {
return "", fmt.Errorf("instance %q lookup failed: %w", instanceName, err)
}
if instance == nil {
return "", fmt.Errorf("instance %q not found for model %q", instanceName, modelName)
}
modelObj, err := m.modelDAO.GetByProviderIDAndInstanceIDAndModelTypeAndModelName(provider.ID, instance.ID, int(modelType), pureModelName)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return "", fmt.Errorf("model %q not found for model type %s", modelName, modelType.String())
}
return "", err
}
if modelObj.Status != "active" {
return "", fmt.Errorf("model %q is disabled", modelName)
}
return modelObj.ID, nil
}
func (m *ModelProviderService) ResolveModelType(tenantID, modelRef string) ([]entity.ModelType, error) {
modelObj, err := m.modelDAO.GetByID(modelRef)
if err == nil {
if modelObj.Status != "active" {
return nil, fmt.Errorf("tenant model id=%s is disabled", modelRef)
}
return []entity.ModelType{entity.ModelType(modelObj.ModelType)}, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
return m.GetModelTypeByName(tenantID, modelRef)
}
// GetModelTypeByName returns the list of model types the given model is enrolled as.
func (m *ModelProviderService) GetModelTypeByName(tenantID, modelName string) ([]entity.ModelType, error) {
common.Debug("GetModelTypeByName",
@@ -2762,7 +2931,7 @@ func (m *ModelProviderService) isImage2TextLLM(tenantID, llmID string) bool {
if m == nil || llmID == "" {
return false
}
modelTypes, err := m.GetModelTypeByName(tenantID, llmID)
modelTypes, err := m.ResolveModelType(tenantID, llmID)
if err != nil {
return false
}
@@ -2786,5 +2955,5 @@ func (m *ModelProviderService) GetChatModelConfig(tenantID string, llmID string)
if m.isImage2TextLLM(tenantID, llmID) {
modelType = entity.ModelTypeImage2Text
}
return m.GetModelConfigFromProviderInstance(tenantID, modelType, llmID)
return m.ResolveModelConfig(tenantID, modelType, llmID)
}

View File

@@ -182,7 +182,7 @@ func TestModelProviderServiceGetModelConfigByID(t *testing.T) {
useModelProviderServiceTestDB(t, db)
seedModelProviderServiceScope(t, db)
driver, modelName, apiConfig, _, err := NewModelProviderService().GetModelConfigByID("user-1", "model-1")
driver, modelName, apiConfig, _, err := NewModelProviderService().GetModelConfigByID("user-1", entity.ModelTypeChat, "model-1")
if err != nil {
t.Fatalf("GetModelConfigByID() error = %v", err)
}

View File

@@ -416,10 +416,46 @@ type ModelItem struct {
ModelProvider *string `json:"model_provider"`
ModelInstance *string `json:"model_instance"`
ModelName *string `json:"model_name"`
ModelID string `json:"model_id"`
ModelType string `json:"model_type"`
Enable bool `json:"enable"`
}
func tenantDefaultModelFields(modelType string) (string, string, entity.ModelType, error) {
switch modelType {
case "chat":
return "llm_id", "tenant_llm_id", entity.ModelTypeChat, nil
case "embedding":
return "embd_id", "tenant_embd_id", entity.ModelTypeEmbedding, nil
case "rerank":
return "rerank_id", "tenant_rerank_id", entity.ModelTypeRerank, nil
case "asr", "speech2text":
return "asr_id", "tenant_asr_id", entity.ModelTypeSpeech2Text, nil
case "vision", "image2text":
return "img2txt_id", "tenant_img2txt_id", entity.ModelTypeImage2Text, nil
case "tts":
return "tts_id", "tenant_tts_id", entity.ModelTypeTTS, nil
case "ocr":
return "ocr_id", "tenant_ocr_id", entity.ModelTypeOCR, nil
default:
return "", "", 0, fmt.Errorf("model type %s is invalid", modelType)
}
}
func ptrStringValue(s *string) string {
if s == nil {
return ""
}
return *s
}
func factoryModelTypeName(modelType string) string {
if modelType == "image2text" {
return "vision"
}
return modelType
}
// GetDefaultModelName returns the full default model ID for a tenant and model type
// Format: modelName@instanceName@providerName or modelName@providerName
// Returns empty string if no default model is set
@@ -459,10 +495,11 @@ func (s *TenantService) GetModelInfo(tenantID string, defaultModel string, model
var providerName *string
var instanceName *string
var modelName *string
if len(defaultChatModelParts) == 3 {
providerName = &defaultChatModelParts[2]
instanceName = &defaultChatModelParts[1]
modelName = &defaultChatModelParts[0]
if len(defaultChatModelParts) >= 3 {
providerName = &defaultChatModelParts[len(defaultChatModelParts)-1]
instanceName = &defaultChatModelParts[len(defaultChatModelParts)-2]
joinedModelName := strings.Join(defaultChatModelParts[:len(defaultChatModelParts)-2], "@")
modelName = &joinedModelName
} else if len(defaultChatModelParts) == 2 {
providerName = &defaultChatModelParts[1]
@@ -495,8 +532,9 @@ func (s *TenantService) GetModelInfo(tenantID string, defaultModel string, model
return nil, nil, nil, false, err
}
if !modelSchema.ModelTypeMap[modelType] {
return nil, nil, nil, false, fmt.Errorf("model %s isn't a chat model", *modelName)
factoryModelType := factoryModelTypeName(modelType)
if !modelSchema.ModelTypeMap[factoryModelType] {
return nil, nil, nil, false, fmt.Errorf("model %s isn't a %s model", *modelName, modelType)
}
var modelEntity *entity.TenantModel
@@ -507,10 +545,14 @@ func (s *TenantService) GetModelInfo(tenantID string, defaultModel string, model
return nil, nil, nil, false, err
}
}
if modelEntity == nil {
return nil, nil, nil, false, fmt.Errorf("model %s isn't available", *modelName)
}
if modelEntity.Status != "active" {
return nil, nil, nil, false, fmt.Errorf("model %s isn't available", *modelName)
}
enable := modelEntity == nil
return providerName, instanceName, modelName, enable, nil
return providerName, instanceName, modelName, true, nil
}
@@ -534,6 +576,7 @@ func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, err
ModelProvider: defaultChatModelProvider,
ModelInstance: defaultChatModelInstance,
ModelName: defaultChatModelName,
ModelID: ptrStringValue(ownedTenant.TenantLLMID),
ModelType: "chat",
Enable: defaultChatModelEnable,
})
@@ -545,6 +588,7 @@ func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, err
ModelProvider: defaultEmbeddingModelProvider,
ModelInstance: defaultEmbeddingModelInstance,
ModelName: defaultEmbeddingModelName,
ModelID: ptrStringValue(ownedTenant.TenantEmbdID),
ModelType: "embedding",
Enable: defaultEmbeddingModelEnable,
})
@@ -556,6 +600,7 @@ func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, err
ModelProvider: defaultRerankModelProvider,
ModelInstance: defaultRerankModelInstance,
ModelName: defaultRerankModelName,
ModelID: ptrStringValue(ownedTenant.TenantRerankID),
ModelType: "rerank",
Enable: defaultRerankModelEnable,
})
@@ -567,6 +612,7 @@ func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, err
ModelProvider: defaultASRModelProvider,
ModelInstance: defaultASRModelInstance,
ModelName: defaultASRModelName,
ModelID: ptrStringValue(ownedTenant.TenantASRID),
ModelType: "asr",
Enable: defaultASREnable,
})
@@ -578,6 +624,7 @@ func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, err
ModelProvider: defaultImage2TextModelProvider,
ModelInstance: defaultImage2TextModelInstance,
ModelName: defaultImage2TextModelName,
ModelID: ptrStringValue(ownedTenant.TenantImg2TxtID),
ModelType: "vision",
Enable: defaultImage2TextModelEnable,
})
@@ -593,6 +640,7 @@ func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, err
ModelProvider: defaultOCRModelProvider,
ModelInstance: defaultOCRModelInstance,
ModelName: defaultOCRModelName,
ModelID: ptrStringValue(ownedTenant.TenantOCRID),
ModelType: "ocr",
Enable: defaultOCRModelEnable,
})
@@ -608,6 +656,7 @@ func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, err
ModelProvider: defaultTTSModelProvider,
ModelInstance: defaultTTSModelInstance,
ModelName: defaultTTSModelName,
ModelID: ptrStringValue(ownedTenant.TenantTTSID),
ModelType: "tts",
Enable: defaultTTSModelEnable,
})
@@ -617,6 +666,11 @@ func (s *TenantService) ListTenantDefaultModels(userID string) ([]ModelItem, err
}
func (s *TenantService) checkModelAvailable(tenantID, providerName, instanceName, modelName, modelType string) error {
_, _, modelTypeBit, err := tenantDefaultModelFields(modelType)
if err != nil {
return err
}
// Check if the provider and instance exists
modelProvider, err := s.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName)
if err != nil {
@@ -633,19 +687,24 @@ func (s *TenantService) checkModelAvailable(tenantID, providerName, instanceName
return err
}
if !modelSchema.ModelTypeMap[modelType] {
return fmt.Errorf("model %s isn't a chat model", modelName)
factoryModelType := factoryModelTypeName(modelType)
if !modelSchema.ModelTypeMap[factoryModelType] {
return fmt.Errorf("model %s isn't a %s model", modelName, modelType)
}
var modelEntity *entity.TenantModel
modelEntity, err = s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(modelProvider.ID, modelInstance.ID, modelName)
if err != nil || modelEntity != nil {
var errString = err.Error()
if errString == "record not found" {
return nil
modelEntity, err := s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(modelProvider.ID, modelInstance.ID, modelName)
if err != nil {
if dao.IsNotFoundErr(err) {
return fmt.Errorf("model %s isn't available", modelName)
}
return err
}
if modelEntity.Status != "active" {
return fmt.Errorf("model %s isn't available", modelName)
}
if !entity.ModelType(modelEntity.ModelType).Has(modelTypeBit) {
return fmt.Errorf("model %s isn't a %s model", modelName, modelType)
}
return nil
}
@@ -662,32 +721,12 @@ func (s *TenantService) SetTenantDefaultModels(userID, modelProvider, modelInsta
ownedTenant := tenantInfos[0]
var defaultModel string
var modelTypeID string
if modelType == "chat" {
modelTypeID = "llm_id"
}
if modelType == "embedding" {
modelTypeID = "embd_id"
}
if modelType == "rerank" {
modelTypeID = "rerank_id"
}
if modelType == "asr" {
modelTypeID = "asr_id"
}
if modelType == "vision" {
modelTypeID = "img2txt_id"
}
if modelType == "tts" {
modelTypeID = "tts_id"
}
if modelType == "ocr" {
modelTypeID = "ocr_id"
}
if modelTypeID == "" {
return fmt.Errorf("model type %s is invalid", modelType)
modelTypeID, tenantModelTypeID, modelTypeBit, err := tenantDefaultModelFields(modelType)
if err != nil {
return err
}
var tenantModelID interface{}
if modelID != "" {
modelEntity, err := s.modelDAO.GetByID(modelID)
if err != nil {
@@ -705,32 +744,50 @@ func (s *TenantService) SetTenantDefaultModels(userID, modelProvider, modelInsta
if providerEntity.TenantID != ownedTenant.TenantID {
return fmt.Errorf("model %s does not belong to your tenant", modelID)
}
if modelEntity.Status != "active" {
return fmt.Errorf("model %s isn't available", modelEntity.ModelName)
}
if !entity.ModelType(modelEntity.ModelType).Has(modelTypeBit) {
return fmt.Errorf("model %s isn't a %s model", modelEntity.ModelName, modelType)
}
if modelProvider == "" {
modelProvider = providerEntity.ProviderName
}
if modelInstance == "" {
modelInstance = instanceEntity.InstanceName
}
if modelName == "" {
modelName = modelEntity.ModelName
}
modelProvider = providerEntity.ProviderName
modelInstance = instanceEntity.InstanceName
modelName = modelEntity.ModelName
tenantModelID = modelID
}
if modelProvider == "" && modelInstance == "" && modelName == "" {
defaultModel = ""
tenantModelID = nil
} else if modelProvider != "" && modelInstance != "" && modelName != "" {
err = s.checkModelAvailable(ownedTenant.TenantID, modelProvider, modelInstance, modelName, modelType)
if err != nil {
return err
}
if modelID == "" {
modelProviderEntity, err := s.modelProviderDAO.GetByTenantIDAndProviderName(ownedTenant.TenantID, modelProvider)
if err != nil {
return err
}
modelInstanceEntity, err := s.modelInstanceDAO.GetByProviderIDAndInstanceName(modelProviderEntity.ID, modelInstance)
if err != nil {
return err
}
modelEntity, err := s.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(modelProviderEntity.ID, modelInstanceEntity.ID, modelName)
if err != nil {
return err
}
tenantModelID = modelEntity.ID
}
defaultModel = fmt.Sprintf("%s@%s@%s", modelName, modelInstance, modelProvider)
} else {
return fmt.Errorf("model provider, instance and name must be specified together")
}
err = s.tenantDAO.Update(ownedTenant.TenantID, map[string]interface{}{
modelTypeID: defaultModel,
modelTypeID: defaultModel,
tenantModelTypeID: tenantModelID,
})
return nil

View File

@@ -250,8 +250,11 @@ func TestSetTenantDefaultModels_WithModelID(t *testing.T) {
t.Fatalf("failed to retrieve tenant: %v", err)
}
expectedDefaultModel := "gpt-4o@default@OpenAI"
expectedDefaultModel := "gpt-4o@dummy@OpenAI"
if tenant.LLMID != expectedDefaultModel {
t.Errorf("expected tenant default LLM to be %q, got %q", expectedDefaultModel, tenant.LLMID)
}
if tenant.TenantLLMID == nil || *tenant.TenantLLMID != modelID {
t.Errorf("expected tenant_llm_id to be %q, got %v", modelID, tenant.TenantLLMID)
}
}