mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-26 01:52:16 +08:00
Go CLI: update list supported models (#15845)
### What problem does this PR solve? Now list supported models will show more info. ``` RAGFlow(api/default)> list supported models from 'gitee' 'test'; +-----------+------------+-------------+----------------------------------------------------------+---------------------------------------------+ | dimension | max_tokens | model_types | name | thinking | +-----------+------------+-------------+----------------------------------------------------------+---------------------------------------------+ | | | | Wan2.7 | | | | | | HappyHorse-1.0 | | | | | | Qwen3.6-27B@Qwen | | | | | | Qwen3.6-35B-A3B@Qwen | | | | 1048576 | [chat] | DeepSeek-V4-Flash@deepseek-ai | map[clear_thinking:true default_value:true] | | | 1048576 | [chat] | DeepSeek-V4-Pro@deepseek-ai | map[clear_thinking:true default_value:true] | +-----------+------------+-------------+----------------------------------------------------------+---------------------------------------------+ ``` ### Type of change - [x] New Feature (non-breaking change which adds functionality) Signed-off-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
@@ -184,7 +184,7 @@ func (m *ModelProviderService) DeleteModelProvider(providerName, userID string)
|
||||
return common.CodeSuccess, nil
|
||||
}
|
||||
|
||||
func (m *ModelProviderService) ListSupportedModels(providerName, instanceName, userID string) ([]string, error) {
|
||||
func (m *ModelProviderService) ListSupportedModels(providerName, instanceName, userID string) ([]map[string]interface{}, error) {
|
||||
|
||||
// Get tenant ID from user
|
||||
tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner")
|
||||
@@ -239,7 +239,22 @@ func (m *ModelProviderService) ListSupportedModels(providerName, instanceName, u
|
||||
}
|
||||
}
|
||||
|
||||
return driver.ListModels(apiConfig)
|
||||
modelList, err := driver.ListModels(apiConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var result []map[string]interface{}
|
||||
for _, model := range modelList {
|
||||
result = append(result, map[string]interface{}{
|
||||
"name": model.Name,
|
||||
"dimension": model.Dimension,
|
||||
"max_tokens": model.MaxTokens,
|
||||
"model_types": model.ModelTypes,
|
||||
"thinking": model.Thinking,
|
||||
})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (m *ModelProviderService) CreateProviderInstance(providerName, instanceName, apiKey, baseURL, region, userID string) (common.ErrorCode, error) {
|
||||
@@ -848,7 +863,7 @@ func (m *ModelProviderService) UpdateModelStatus(providerName, instanceName, mod
|
||||
return common.CodeServerError, errors.New("fail to get UUID")
|
||||
}
|
||||
|
||||
var modelSchema *entity.Model
|
||||
var modelSchema *modelModule.Model
|
||||
modelSchema, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName))
|
||||
@@ -920,7 +935,7 @@ func (m *ModelProviderService) ChatToModelWithMessages(providerName, instanceNam
|
||||
return nil, common.CodeNotFound, errors.New("provider not found")
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName))
|
||||
@@ -1028,7 +1043,7 @@ func (m *ModelProviderService) ChatToModelStreamWithSender(providerName, instanc
|
||||
return common.CodeNotFound, err
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return common.CodeNotFound, err
|
||||
@@ -1132,7 +1147,7 @@ func (m *ModelProviderService) EmbedText(providerName, instanceName, modelName,
|
||||
return nil, common.CodeNotFound, errors.New("provider not found")
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName))
|
||||
@@ -1243,7 +1258,7 @@ func (m *ModelProviderService) RerankDocument(providerName, instanceName, modelN
|
||||
return nil, common.CodeNotFound, errors.New("provider not found")
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName))
|
||||
@@ -1348,7 +1363,7 @@ func (m *ModelProviderService) TranscribeAudio(providerName, instanceName, model
|
||||
return nil, common.CodeNotFound, errors.New("provider not found")
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName))
|
||||
@@ -1452,7 +1467,7 @@ func (m *ModelProviderService) TranscribeAudioStream(providerName, instanceName,
|
||||
return common.CodeNotFound, err
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return common.CodeNotFound, err
|
||||
@@ -1553,7 +1568,7 @@ func (m *ModelProviderService) AudioSpeech(providerName, instanceName, modelName
|
||||
return nil, common.CodeNotFound, errors.New("provider not found")
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName))
|
||||
@@ -1656,7 +1671,7 @@ func (m *ModelProviderService) AudioSpeechStream(providerName, instanceName, mod
|
||||
return common.CodeNotFound, err
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return common.CodeNotFound, err
|
||||
@@ -1757,7 +1772,7 @@ func (m *ModelProviderService) OCRFile(providerName, instanceName, modelName, us
|
||||
return nil, common.CodeNotFound, errors.New("provider not found")
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName))
|
||||
@@ -1867,7 +1882,7 @@ func (m *ModelProviderService) ParseFile(providerName, instanceName, modelName,
|
||||
return nil, common.CodeNotFound, errors.New("provider not found")
|
||||
}
|
||||
|
||||
var model *entity.Model = nil
|
||||
var model *modelModule.Model = nil
|
||||
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
if err != nil {
|
||||
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName))
|
||||
@@ -2218,7 +2233,11 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin
|
||||
apiConfig := &modelModule.APIConfig{ApiKey: &apiKey, Region: ®ion}
|
||||
maxTokens := 0
|
||||
if mi, _ := dao.GetModelProviderManager().GetModelByName("Builtin", pureModelName); mi != nil {
|
||||
maxTokens = mi.MaxTokens
|
||||
if mi.MaxTokens == nil {
|
||||
maxTokens = 0
|
||||
} else {
|
||||
maxTokens = *mi.MaxTokens
|
||||
}
|
||||
}
|
||||
return builtinDriver, pureModelName, apiConfig, maxTokens, nil
|
||||
}
|
||||
@@ -2275,7 +2294,11 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin
|
||||
}
|
||||
maxTokens := 0
|
||||
if mi, _ := dao.GetModelProviderManager().GetModelByName(providerName, pureModelName); mi != nil {
|
||||
maxTokens = mi.MaxTokens
|
||||
if mi.MaxTokens == nil {
|
||||
maxTokens = 0
|
||||
} else {
|
||||
maxTokens = *mi.MaxTokens
|
||||
}
|
||||
}
|
||||
apiConfig := &modelModule.APIConfig{ApiKey: &apiKey, Region: ®ion, BaseURL: &baseURL}
|
||||
return driver, modelObj.ModelName, apiConfig, maxTokens, nil
|
||||
@@ -2309,7 +2332,7 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin
|
||||
if targetProvider == nil {
|
||||
return nil, "", nil, 0, fmt.Errorf("model provider config not found: %s", providerName)
|
||||
}
|
||||
var llmInfo *entity.Model
|
||||
var llmInfo *modelModule.Model
|
||||
for i := range targetProvider.Models {
|
||||
if strings.EqualFold(targetProvider.Models[i].Name, pureModelName) {
|
||||
llmInfo = targetProvider.Models[i]
|
||||
@@ -2324,7 +2347,11 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin
|
||||
return nil, "", nil, 0, driverErr
|
||||
}
|
||||
apiConfig := &modelModule.APIConfig{ApiKey: &apiKey, Region: ®ion, BaseURL: &baseURL}
|
||||
return driver, llmInfo.Name, apiConfig, llmInfo.MaxTokens, nil
|
||||
maxTokens := 0
|
||||
if llmInfo.MaxTokens != nil {
|
||||
maxTokens = *llmInfo.MaxTokens
|
||||
}
|
||||
return driver, llmInfo.Name, apiConfig, maxTokens, nil
|
||||
}
|
||||
|
||||
// getModelConfig returns the model driver, model name, API config, and max tokens for a model
|
||||
@@ -2381,7 +2408,11 @@ func (m *ModelProviderService) getModelConfig(tenantID, compositeModelName strin
|
||||
modelInfo, err := dao.GetModelProviderManager().GetModelByName(providerName, modelName)
|
||||
maxTokens := 0
|
||||
if err == nil && modelInfo != nil {
|
||||
maxTokens = modelInfo.MaxTokens
|
||||
if modelInfo.MaxTokens == nil {
|
||||
maxTokens = 0
|
||||
} else {
|
||||
maxTokens = *modelInfo.MaxTokens
|
||||
}
|
||||
}
|
||||
|
||||
// For Builtin provider, use empty APIKey and skip tenant_model lookup
|
||||
@@ -2423,6 +2454,6 @@ func (m *ModelProviderService) ListAllModels(pageIndex, pageSize int) ([]map[str
|
||||
return models, nil
|
||||
}
|
||||
|
||||
func (m *ModelProviderService) ShowModel(modelName string) (*entity.Model, error) {
|
||||
func (m *ModelProviderService) ShowModel(modelName string) (*modelModule.Model, error) {
|
||||
return dao.GetModelProviderManager().GetModelByNameOrAlias(modelName), nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user