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:
Jin Hai
2026-06-09 19:01:00 +08:00
committed by GitHub
parent f0efa63bf2
commit 719ce15c95
68 changed files with 357 additions and 258 deletions
+10 -6
View File
@@ -21,6 +21,7 @@ import (
"log" "log"
"ragflow/internal/common" "ragflow/internal/common"
"ragflow/internal/entity" "ragflow/internal/entity"
"ragflow/internal/entity/models"
"strings" "strings"
"time" "time"
@@ -34,7 +35,7 @@ import (
) )
var DB *gorm.DB var DB *gorm.DB
var modelProviderManager *entity.ProviderManager var modelProviderManager *models.ProviderManager
// LLMFactoryConfig represents a single LLM factory configuration // LLMFactoryConfig represents a single LLM factory configuration
type LLMFactoryConfig struct { type LLMFactoryConfig struct {
@@ -106,8 +107,8 @@ func InitDB() error {
sqlDB.SetMaxOpenConns(100) sqlDB.SetMaxOpenConns(100)
sqlDB.SetConnMaxLifetime(time.Hour) sqlDB.SetConnMaxLifetime(time.Hour)
// Auto migrate all models // Auto migrate all dataModels
models := []interface{}{ dataModels := []interface{}{
&entity.User{}, &entity.User{},
&entity.Tenant{}, &entity.Tenant{},
&entity.UserTenant{}, &entity.UserTenant{},
@@ -150,7 +151,7 @@ func InitDB() error {
&entity.TenantModelGroup{}, &entity.TenantModelGroup{},
} }
for _, m := range models { for _, m := range dataModels {
if err = autoMigrateSafely(DB, m); err != nil { if err = autoMigrateSafely(DB, m); err != nil {
return fmt.Errorf("failed to migrate model %T: %w", m, err) return fmt.Errorf("failed to migrate model %T: %w", m, err)
} }
@@ -163,11 +164,14 @@ func InitDB() error {
common.Info("Database connected and migrated successfully") common.Info("Database connected and migrated successfully")
modelProviderManager, err = entity.NewProviderManager("conf/models") err = models.InitProviderManager("conf/models")
if err != nil { if err != nil {
log.Fatal("Failed to load model providers:", err) log.Fatal("Failed to load model providers:", err)
} }
modelProviderManager = models.GetProviderManager()
common.Info("Model providers loaded successfully") common.Info("Model providers loaded successfully")
return nil return nil
} }
@@ -177,7 +181,7 @@ func GetDB() *gorm.DB {
} }
// GetModelProviderManager get database instance // GetModelProviderManager get database instance
func GetModelProviderManager() *entity.ProviderManager { func GetModelProviderManager() *models.ProviderManager {
return modelProviderManager return modelProviderManager
} }
+5 -3
View File
@@ -877,7 +877,7 @@ func (a *AI302Model) ParseFile(modelName *string, content []byte, documentURL *s
}, nil }, nil
} }
func (a *AI302Model) ListModels(apiConfig *APIConfig) ([]string, error) { func (a *AI302Model) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := a.baseModel.APIConfigCheck(apiConfig); err != nil { if err := a.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -927,12 +927,14 @@ func (a *AI302Model) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("models response missing data") return nil, fmt.Errorf("models response missing data")
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if strings.TrimSpace(model.ID) == "" { if strings.TrimSpace(model.ID) == "" {
return nil, fmt.Errorf("models response contains empty id") return nil, fmt.Errorf("models response contains empty id")
} }
models = append(models, strings.TrimSpace(model.ID)) models = append(models, ListModelResponse{
Name: strings.TrimSpace(model.ID),
})
} }
return models, nil return models, nil
+5 -3
View File
@@ -590,7 +590,7 @@ type AliyunModelList struct {
Output AliyunModelOutput `json:"output"` Output AliyunModelOutput `json:"output"`
} }
func (a *AliyunModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (a *AliyunModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := a.baseModel.APIConfigCheck(apiConfig); err != nil { if err := a.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -643,10 +643,12 @@ func (a *AliyunModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
var models []string var models []ListModelResponse
for _, model := range modelList.Output.Models { for _, model := range modelList.Output.Models {
modelName := model.ModelName modelName := model.ModelName
models = append(models, modelName) models = append(models, ListModelResponse{
Name: modelName,
})
} }
return models, nil return models, nil
+5 -3
View File
@@ -376,7 +376,7 @@ func parseAnthropicChatResponse(body []byte) (string, string, error) {
return answer.String(), reasoning.String(), nil return answer.String(), reasoning.String(), nil
} }
func (a *AnthropicModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (a *AnthropicModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := a.baseModel.APIConfigCheck(apiConfig); err != nil { if err := a.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -425,10 +425,12 @@ func (a *AnthropicModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if err = json.Unmarshal(body, &result); err != nil { if err = json.Unmarshal(body, &result); err != nil {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, item := range result.Data { for _, item := range result.Data {
if item.ID != "" { if item.ID != "" {
models = append(models, item.ID) models = append(models, ListModelResponse{
Name: item.ID,
})
} }
} }
return models, nil return models, nil
+9 -5
View File
@@ -328,7 +328,7 @@ func (a *AstraflowModel) ChatStreamlyWithSender(modelName string, messages []Mes
return nil return nil
} }
func (a *AstraflowModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (a *AstraflowModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := a.baseModel.APIConfigCheck(apiConfig); err != nil { if err := a.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -373,17 +373,21 @@ func (a *AstraflowModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0, len(data)) models := make([]ListModelResponse, 0, len(data))
var modelMap map[string]interface{}
for _, m := range data { for _, m := range data {
modelMap, ok := m.(map[string]interface{}) modelMap, ok = m.(map[string]interface{})
if !ok { if !ok {
continue continue
} }
id, ok := modelMap["id"].(string) var id string
id, ok = modelMap["id"].(string)
if !ok { if !ok {
continue continue
} }
models = append(models, id) models = append(models, ListModelResponse{
Name: id,
})
} }
return models, nil return models, nil
} }
+5 -3
View File
@@ -299,7 +299,7 @@ type avianModelInfo struct {
ID string `json:"id"` ID string `json:"id"`
} }
func (a *AvianModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (a *AvianModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := a.baseModel.APIConfigCheck(apiConfig); err != nil { if err := a.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -340,10 +340,12 @@ func (a *AvianModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(result)) models := make([]ListModelResponse, 0, len(result))
for _, model := range result { for _, model := range result {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{
Name: model.ID,
})
} }
} }
return models, nil return models, nil
+5 -3
View File
@@ -424,7 +424,7 @@ func (a *AzureOpenAIModel) Embed(modelName *string, texts []string, apiConfig *A
} }
// ListModels returns the deployment names visible to the configured API key. // ListModels returns the deployment names visible to the configured API key.
func (a *AzureOpenAIModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (a *AzureOpenAIModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := a.baseModel.APIConfigCheck(apiConfig); err != nil { if err := a.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -472,7 +472,7 @@ func (a *AzureOpenAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid deployments list format") return nil, fmt.Errorf("invalid deployments list format")
} }
models := make([]string, 0, len(data)) models := make([]ListModelResponse, 0, len(data))
for _, item := range data { for _, item := range data {
m, ok := item.(map[string]interface{}) m, ok := item.(map[string]interface{})
if !ok { if !ok {
@@ -483,7 +483,9 @@ func (a *AzureOpenAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, id) models = append(models, ListModelResponse{
Name: id,
})
} }
return models, nil return models, nil
+1 -1
View File
@@ -426,7 +426,7 @@ func (b *BaichuanModel) ParseFile(modelName *string, content []byte, url *string
return nil, fmt.Errorf("%s, no such method", b.Name()) return nil, fmt.Errorf("%s, no such method", b.Name())
} }
func (b *BaichuanModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (b *BaichuanModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("no such method") return nil, fmt.Errorf("no such method")
} }
+5 -3
View File
@@ -744,7 +744,7 @@ func (b *BaiduModel) OCRFile(modelName *string, content []byte, fileURL *string,
return &ocrResponse, nil return &ocrResponse, nil
} }
func (b *BaiduModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (b *BaiduModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := b.baseModel.APIConfigCheck(apiConfig); err != nil { if err := b.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -795,11 +795,13 @@ func (b *BaiduModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
// convert result["data"] to []map[string]interface{} // convert result["data"] to []map[string]interface{}
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["data"].([]interface{}) { for _, model := range result["data"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["id"].(string) modelName := modelMap["id"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{
Name: modelName,
})
} }
return models, nil return models, nil
+7 -5
View File
@@ -777,7 +777,7 @@ type bedrockListModelsResponse struct {
// configured credentials. The control plane lives at // configured credentials. The control plane lives at
// bedrock.{region}.amazonaws.com (not bedrock-runtime), signs against // bedrock.{region}.amazonaws.com (not bedrock-runtime), signs against
// the "bedrock" service, and is GET-only. // the "bedrock" service, and is GET-only.
func (b *BedrockModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (b *BedrockModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := b.baseModel.APIConfigCheck(apiConfig); err != nil { if err := b.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -805,7 +805,7 @@ func (b *BedrockModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("bedrock: build request: %w", err) return nil, fmt.Errorf("bedrock: build request: %w", err)
} }
req.Header.Set("Accept", "application/json") req.Header.Set("Accept", "application/json")
if err := signBedrockRequest(ctx, req, nil, creds, bedrockControlService, region); err != nil { if err = signBedrockRequest(ctx, req, nil, creds, bedrockControlService, region); err != nil {
return nil, err return nil, err
} }
@@ -824,15 +824,17 @@ func (b *BedrockModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
var parsed bedrockListModelsResponse var parsed bedrockListModelsResponse
if err := json.Unmarshal(respBody, &parsed); err != nil { if err = json.Unmarshal(respBody, &parsed); err != nil {
return nil, fmt.Errorf("bedrock: parse ListModels response: %w", err) return nil, fmt.Errorf("bedrock: parse ListModels response: %w", err)
} }
models := make([]string, 0, len(parsed.ModelSummaries)) models := make([]ListModelResponse, 0, len(parsed.ModelSummaries))
for _, m := range parsed.ModelSummaries { for _, m := range parsed.ModelSummaries {
if m.ModelID == "" { if m.ModelID == "" {
continue continue
} }
models = append(models, m.ModelID) models = append(models, ListModelResponse{
Name: m.ModelID,
})
} }
return models, nil return models, nil
} }
+7 -3
View File
@@ -115,7 +115,7 @@ func (b *BuiltinModel) Embed(modelName *string, texts []string, apiConfig *APICo
for i, emb := range embeddings { for i, emb := range embeddings {
result[i] = EmbeddingData{ result[i] = EmbeddingData{
Embedding: emb, Embedding: emb,
Index: i, Index: i,
} }
} }
@@ -150,8 +150,12 @@ func (b *BuiltinModel) ParseFile(modelName *string, content []byte, url *string,
return nil, fmt.Errorf("builtin model does not support parse file") return nil, fmt.Errorf("builtin model does not support parse file")
} }
func (b *BuiltinModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (b *BuiltinModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return []string{b.model}, nil return []ListModelResponse{
{
Name: b.model,
},
}, nil
} }
func (b *BuiltinModel) Balance(apiConfig *APIConfig) (map[string]interface{}, error) { func (b *BuiltinModel) Balance(apiConfig *APIConfig) (map[string]interface{}, error) {
+5 -3
View File
@@ -650,7 +650,7 @@ func (c *CoHereModel) ParseFile(modelName *string, content []byte, url *string,
return nil, fmt.Errorf("%s, no such method", c.Name()) return nil, fmt.Errorf("%s, no such method", c.Name())
} }
func (c *CoHereModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (c *CoHereModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := c.baseModel.APIConfigCheck(apiConfig); err != nil { if err := c.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -692,12 +692,14 @@ func (c *CoHereModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
if modelsRaw, ok := result["models"].([]interface{}); ok { if modelsRaw, ok := result["models"].([]interface{}); ok {
for _, model := range modelsRaw { for _, model := range modelsRaw {
if modelMap, ok := model.(map[string]interface{}); ok { if modelMap, ok := model.(map[string]interface{}); ok {
if modelName, ok := modelMap["name"].(string); ok { if modelName, ok := modelMap["name"].(string); ok {
models = append(models, modelName) models = append(models, ListModelResponse{
Name: modelName,
})
} }
} }
} }
+6 -4
View File
@@ -244,16 +244,18 @@ type cometapiModelCatalogItem struct {
ID string `json:"id"` ID string `json:"id"`
} }
func parseCometAPIModelCatalog(body []byte) ([]string, error) { func parseCometAPIModelCatalog(body []byte) ([]ListModelResponse, error) {
var parsed cometapiModelCatalogResponse var parsed cometapiModelCatalogResponse
if err := json.Unmarshal(body, &parsed); err != nil { if err := json.Unmarshal(body, &parsed); err != nil {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(parsed.Data)) models := make([]ListModelResponse, 0, len(parsed.Data))
for _, model := range parsed.Data { for _, model := range parsed.Data {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{
Name: model.ID,
})
} }
} }
return models, nil return models, nil
@@ -498,7 +500,7 @@ func (c *CometAPIModel) Embed(modelName *string, texts []string, apiConfig *APIC
} }
// ListModels returns the public CometAPI model catalog. // ListModels returns the public CometAPI model catalog.
func (c *CometAPIModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (c *CometAPIModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
url, err := c.endpointURL(cometapiRegion(apiConfig), c.baseModel.URLSuffix.Models) url, err := c.endpointURL(cometapiRegion(apiConfig), c.baseModel.URLSuffix.Models)
if err != nil { if err != nil {
return nil, err return nil, err
+5 -3
View File
@@ -840,7 +840,7 @@ func (d *DeepInfraModel) ParseFile(modelName *string, content []byte, url *strin
return nil, fmt.Errorf("%s no such method", d.Name()) return nil, fmt.Errorf("%s no such method", d.Name())
} }
func (d *DeepInfraModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (d *DeepInfraModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
resolvedBaseURL, err := d.baseModel.GetBaseURL(apiConfig) resolvedBaseURL, err := d.baseModel.GetBaseURL(apiConfig)
if err != nil { if err != nil {
@@ -888,10 +888,12 @@ func (d *DeepInfraModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to unmarshal response: %w", err) return nil, fmt.Errorf("failed to unmarshal response: %w", err)
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result { for _, model := range result {
if model.ModelName != "" { if model.ModelName != "" {
models = append(models, model.ModelName) models = append(models, ListModelResponse{
Name: model.ModelName,
})
} }
} }
+5 -3
View File
@@ -442,7 +442,7 @@ type DSModelList struct {
Models []DSModel `json:"data"` Models []DSModel `json:"data"`
} }
func (d *DeepSeekModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (d *DeepSeekModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := d.baseModel.APIConfigCheck(apiConfig); err != nil { if err := d.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -493,9 +493,11 @@ func (d *DeepSeekModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
var models []string var models []ListModelResponse
for _, model := range modelList.Models { for _, model := range modelList.Models {
models = append(models, model.ID) models = append(models, ListModelResponse{
Name: model.ID,
})
} }
return models, nil return models, nil
+1 -1
View File
@@ -58,7 +58,7 @@ func (d *DummyModel) Embed(modelName *string, texts []string, apiConfig *APIConf
return nil, fmt.Errorf("not implemented") return nil, fmt.Errorf("not implemented")
} }
func (d *DummyModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (d *DummyModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("not implemented") return nil, fmt.Errorf("not implemented")
} }
+5 -3
View File
@@ -352,7 +352,7 @@ func (f *FishAudioModel) ParseFile(modelName *string, content []byte, url *strin
return nil, fmt.Errorf("%s, no such method", f.Name()) return nil, fmt.Errorf("%s, no such method", f.Name())
} }
func (f *FishAudioModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (f *FishAudioModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := f.baseModel.APIConfigCheck(apiConfig); err != nil { if err := f.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -399,9 +399,11 @@ func (f *FishAudioModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(result.Items)) models := make([]ListModelResponse, 0, len(result.Items))
for _, item := range result.Items { for _, item := range result.Items {
models = append(models, item.Title) models = append(models, ListModelResponse{
Name: item.Title,
})
} }
return models, nil return models, nil
+1 -1
View File
@@ -358,7 +358,7 @@ func (f *FuturMixModel) Rerank(modelName *string, query string, documents []stri
} }
// ListModels is not documented as a public endpoint by FuturMix. // ListModels is not documented as a public endpoint by FuturMix.
func (f *FuturMixModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (f *FuturMixModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("%s, no such method", f.Name()) return nil, fmt.Errorf("%s, no such method", f.Name())
} }
+14 -3
View File
@@ -852,7 +852,7 @@ func (g *GiteeModel) getParseFile(baseURL *string, apiKey, taskID *string, timeO
return nil, nil return nil, nil
} }
func (g *GiteeModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (g *GiteeModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
resolvedBaseURL, err := g.baseModel.GetBaseURL(apiConfig) resolvedBaseURL, err := g.baseModel.GetBaseURL(apiConfig)
if err != nil { if err != nil {
@@ -899,13 +899,24 @@ func (g *GiteeModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
var models []string var models []ListModelResponse
for _, model := range modelList.Models { for _, model := range modelList.Models {
modelName := model.ID modelName := model.ID
var modelResponse ListModelResponse
pm := GetProviderManager()
modelEntity := pm.GetModelByNameOrAlias(modelName)
if model.OwnedBy != "" { if model.OwnedBy != "" {
modelName = model.ID + "@" + model.OwnedBy modelName = model.ID + "@" + model.OwnedBy
} }
models = append(models, modelName) modelResponse.Name = modelName
if modelEntity != nil {
modelResponse.Dimension = modelEntity.Dimension
modelResponse.MaxTokens = modelEntity.MaxTokens
modelResponse.ModelTypes = modelEntity.ModelTypes
modelResponse.Thinking = modelEntity.Thinking
}
models = append(models, modelResponse)
} }
return models, nil return models, nil
+9 -5
View File
@@ -30,8 +30,8 @@ type googleModelPage struct {
nextPageToken string nextPageToken string
} }
func collectGoogleModelNames(ctx context.Context, listPage func(context.Context, string) (googleModelPage, error)) ([]string, error) { func collectGoogleModelNames(ctx context.Context, listPage func(context.Context, string) (googleModelPage, error)) ([]ListModelResponse, error) {
var modelNames []string var modelNames []ListModelResponse
pageToken := "" pageToken := ""
for { for {
@@ -40,7 +40,11 @@ func collectGoogleModelNames(ctx context.Context, listPage func(context.Context,
return nil, err return nil, err
} }
modelNames = append(modelNames, page.items...) for _, modelName := range page.items {
modelNames = append(modelNames, ListModelResponse{
Name: modelName,
})
}
if page.nextPageToken == "" { if page.nextPageToken == "" {
return modelNames, nil return modelNames, nil
} }
@@ -48,7 +52,7 @@ func collectGoogleModelNames(ctx context.Context, listPage func(context.Context,
} }
} }
var googleListModels = func(ctx context.Context, config *genai.ClientConfig) ([]string, error) { var googleListModels = func(ctx context.Context, config *genai.ClientConfig) ([]ListModelResponse, error) {
client, err := genai.NewClient(ctx, config) client, err := genai.NewClient(ctx, config)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -339,7 +343,7 @@ func (g *GoogleModel) Embed(modelName *string, texts []string, apiConfig *APICon
return result, nil return result, nil
} }
func (g *GoogleModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (g *GoogleModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := g.baseModel.APIConfigCheck(apiConfig); err != nil { if err := g.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
+5 -3
View File
@@ -335,7 +335,7 @@ type gpustackModelsResponse struct {
Data []gpustackModelInfo `json:"data"` Data []gpustackModelInfo `json:"data"`
} }
func (g *GPUStackModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (g *GPUStackModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := g.baseModel.APIConfigCheck(apiConfig); err != nil { if err := g.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -378,10 +378,12 @@ func (g *GPUStackModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(parsed.Data)) models := make([]ListModelResponse, 0, len(parsed.Data))
for _, m := range parsed.Data { for _, m := range parsed.Data {
if m.ID != "" { if m.ID != "" {
models = append(models, m.ID) models = append(models, ListModelResponse{
Name: m.ID,
})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -327,7 +327,7 @@ type groqListModelsResponse struct {
Error interface{} `json:"error"` Error interface{} `json:"error"`
} }
func (g *GroqModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (g *GroqModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := g.baseModel.APIConfigCheck(apiConfig); err != nil { if err := g.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -369,10 +369,10 @@ func (g *GroqModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("groq: upstream error: %v", result.Error) return nil, fmt.Errorf("groq: upstream error: %v", result.Error)
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{Name: model.ID})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -631,7 +631,7 @@ func (h *HuaweiCloudModel) ParseFile(modelName *string, content []byte, url *str
return nil, fmt.Errorf("%s, no such method", h.Name()) return nil, fmt.Errorf("%s, no such method", h.Name())
} }
func (h *HuaweiCloudModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (h *HuaweiCloudModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := h.baseModel.APIConfigCheck(apiConfig); err != nil { if err := h.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -680,10 +680,10 @@ func (h *HuaweiCloudModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to decode response: %w", err) return nil, fmt.Errorf("failed to decode response: %w", err)
} }
models := make([]string, 0, len(parsed.Data)) models := make([]ListModelResponse, 0, len(parsed.Data))
for _, item := range parsed.Data { for _, item := range parsed.Data {
if item.ID != "" { if item.ID != "" {
models = append(models, item.ID) models = append(models, ListModelResponse{Name: item.ID})
} }
} }
if len(models) == 0 { if len(models) == 0 {
+3 -3
View File
@@ -460,7 +460,7 @@ func (h *HuggingFaceModel) ParseFile(modelName *string, content []byte, url *str
return nil, fmt.Errorf("%s, no such method", h.Name()) return nil, fmt.Errorf("%s, no such method", h.Name())
} }
func (h *HuggingFaceModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (h *HuggingFaceModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := h.baseModel.APIConfigCheck(apiConfig); err != nil { if err := h.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -511,11 +511,11 @@ func (h *HuggingFaceModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["data"].([]interface{}) { for _, model := range result["data"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["id"].(string) modelName := modelMap["id"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -310,7 +310,7 @@ func (h *HunyuanModel) ChatStreamlyWithSender(modelName string, messages []Messa
return nil return nil
} }
func (h *HunyuanModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (h *HunyuanModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := h.baseModel.APIConfigCheck(apiConfig); err != nil { if err := h.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -355,7 +355,7 @@ func (h *HunyuanModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0, len(data)) models := make([]ListModelResponse, 0, len(data))
for _, m := range data { for _, m := range data {
modelMap, ok := m.(map[string]interface{}) modelMap, ok := m.(map[string]interface{})
if !ok { if !ok {
@@ -365,7 +365,7 @@ func (h *HunyuanModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, id) models = append(models, ListModelResponse{Name: id})
} }
return models, nil return models, nil
} }
+3 -3
View File
@@ -545,7 +545,7 @@ func (j *JieKouAIModel) ParseFile(modelName *string, content []byte, url *string
return nil, fmt.Errorf("%s, no such method", j.Name()) return nil, fmt.Errorf("%s, no such method", j.Name())
} }
func (j *JieKouAIModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (j *JieKouAIModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := j.baseModel.APIConfigCheck(apiConfig); err != nil { if err := j.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -592,12 +592,12 @@ func (j *JieKouAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("models response missing data") return nil, fmt.Errorf("models response missing data")
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if strings.TrimSpace(model.ID) == "" { if strings.TrimSpace(model.ID) == "" {
return nil, fmt.Errorf("models response contains empty id") return nil, fmt.Errorf("models response contains empty id")
} }
models = append(models, strings.TrimSpace(model.ID)) models = append(models, ListModelResponse{Name: strings.TrimSpace(model.ID)})
} }
return models, nil return models, nil
+3 -3
View File
@@ -318,7 +318,7 @@ func (j *JinaModel) Rerank(modelName *string, query string, documents []string,
return &rerankResponse, nil return &rerankResponse, nil
} }
func (j *JinaModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (j *JinaModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
resolvedBaseURL, err := j.baseModel.GetBaseURL(apiConfig) resolvedBaseURL, err := j.baseModel.GetBaseURL(apiConfig)
if err != nil { if err != nil {
@@ -355,11 +355,11 @@ func (j *JinaModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
// convert result["data"] to []map[string]interface{} // convert result["data"] to []map[string]interface{}
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["data"].([]interface{}) { for _, model := range result["data"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["name"].(string) modelName := modelMap["name"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -500,7 +500,7 @@ func (l *LmStudioModel) ParseFile(modelName *string, content []byte, url *string
} }
// ListModels list supported models // ListModels list supported models
func (l *LmStudioModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (l *LmStudioModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := l.baseModel.APIConfigCheck(apiConfig); err != nil { if err := l.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -561,11 +561,11 @@ func (l *LmStudioModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
// convert result["data"] 2 []map[string]interface{} // convert result["data"] 2 []map[string]interface{}
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["data"].([]interface{}) { for _, model := range result["data"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["id"].(string) modelName := modelMap["id"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -583,7 +583,7 @@ func (l *LocalAIModel) Rerank(modelName *string, query string, documents []strin
return rerankResponse, nil return rerankResponse, nil
} }
func (l *LocalAIModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (l *LocalAIModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := l.baseModel.APIConfigCheck(apiConfig); err != nil { if err := l.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -632,7 +632,7 @@ func (l *LocalAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -642,7 +642,7 @@ func (l *LocalAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -348,7 +348,7 @@ type longCatListModelsResponse struct {
const longCatMaxListModelsResponseBytes = 1 << 20 const longCatMaxListModelsResponseBytes = 1 << 20
func (l *LongCatModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (l *LongCatModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := l.baseModel.APIConfigCheck(apiConfig); err != nil { if err := l.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -398,10 +398,10 @@ func (l *LongCatModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{Name: model.ID})
} }
} }
return models, nil return models, nil
+1 -1
View File
@@ -91,7 +91,7 @@ func (m *MinerUModel) OCRFile(modelName *string, content []byte, url *string, ap
return nil, fmt.Errorf("%s no such method", m.Name()) return nil, fmt.Errorf("%s no such method", m.Name())
} }
func (m *MinerUModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (m *MinerUModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("%s no such method", m.Name()) return nil, fmt.Errorf("%s no such method", m.Name())
} }
+1 -1
View File
@@ -93,7 +93,7 @@ func (m *MinerULocalModel) OCRFile(modelName *string, content []byte, url *strin
return nil, fmt.Errorf("%s no such method", m.Name()) return nil, fmt.Errorf("%s no such method", m.Name())
} }
func (m *MinerULocalModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (m *MinerULocalModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("%s no such method", m.Name()) return nil, fmt.Errorf("%s no such method", m.Name())
} }
+3 -3
View File
@@ -388,7 +388,7 @@ func (m *MinimaxModel) Embed(modelName *string, texts []string, apiConfig *APICo
return nil, fmt.Errorf("not implemented") return nil, fmt.Errorf("not implemented")
} }
func (m *MinimaxModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (m *MinimaxModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := m.baseModel.APIConfigCheck(apiConfig); err != nil { if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -438,12 +438,12 @@ func (m *MinimaxModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("models response missing data") return nil, fmt.Errorf("models response missing data")
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if strings.TrimSpace(model.ID) == "" { if strings.TrimSpace(model.ID) == "" {
return nil, fmt.Errorf("models response contains empty id") return nil, fmt.Errorf("models response contains empty id")
} }
models = append(models, strings.TrimSpace(model.ID)) models = append(models, ListModelResponse{Name: strings.TrimSpace(model.ID)})
} }
return models, nil return models, nil
+3 -3
View File
@@ -461,7 +461,7 @@ func (m *MistralModel) Embed(modelName *string, texts []string, apiConfig *APICo
} }
// ListModels returns the list of model ids visible to the API key. // ListModels returns the list of model ids visible to the API key.
func (m *MistralModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (m *MistralModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := m.baseModel.APIConfigCheck(apiConfig); err != nil { if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -508,7 +508,7 @@ func (m *MistralModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -518,7 +518,7 @@ func (m *MistralModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
@@ -14,7 +14,7 @@
// limitations under the License. // limitations under the License.
// //
package entity package models
import ( import (
"bytes" "bytes"
@@ -22,7 +22,6 @@ import (
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"ragflow/internal/entity/models"
"strings" "strings"
) )
@@ -157,10 +156,11 @@ type ModelThinking struct {
// Model represents a single LLM model // Model represents a single LLM model
type Model struct { type Model struct {
Name string `json:"name"` Name string `json:"name"`
MaxTokens int `json:"max_tokens"` MaxTokens *int `json:"max_tokens"`
ModelTypes []string `json:"model_types"` ModelTypes []string `json:"model_types"`
Thinking *ModelThinking `json:"thinking"` Thinking *ModelThinking `json:"thinking"`
Class *string `json:"class"` Class *string `json:"class"`
Dimension *int `json:"dimension"` // used by embedding models
Alias []string `json:"alias"` Alias []string `json:"alias"`
ModelTypeMap map[string]bool ModelTypeMap map[string]bool
} }
@@ -169,11 +169,11 @@ type Model struct {
type Provider struct { type Provider struct {
Name string `json:"name"` Name string `json:"name"`
URL map[string]string `json:"url"` URL map[string]string `json:"url"`
URLSuffix models.URLSuffix `json:"url_suffix"` URLSuffix URLSuffix `json:"url_suffix"`
Models []*Model `json:"models"` Models []*Model `json:"models"`
Features Features `json:"features"` Features Features `json:"features"`
Class string `json:"class"` Class string `json:"class"`
ModelDriver models.ModelDriver ModelDriver ModelDriver
} }
// ProviderManager manages provider and model operations // ProviderManager manages provider and model operations
@@ -215,17 +215,23 @@ func decodeProviderConfig(data []byte) (Provider, error) {
return provider, nil return provider, nil
} }
// NewProviderManager creates a new ProviderManager by reading all JSON files from a directory var providerManager *ProviderManager
func NewProviderManager(dirPath string) (*ProviderManager, error) {
func GetProviderManager() *ProviderManager {
return providerManager
}
// InitProviderManager creates a new ProviderManager by reading all JSON files from a directory
func InitProviderManager(dirPath string) error {
providers := []Provider{} providers := []Provider{}
// Read all files in the directory // Read all files in the directory
files, err := os.ReadDir(dirPath) files, err := os.ReadDir(dirPath)
if err != nil { if err != nil {
return nil, fmt.Errorf("error reading directory %s: %w", dirPath, err) return fmt.Errorf("error reading directory %s: %w", dirPath, err)
} }
modelFactory := models.NewModelFactory() modelFactory := NewModelFactory()
// Iterate through all files // Iterate through all files
for _, file := range files { for _, file := range files {
@@ -246,13 +252,13 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) {
var data []byte var data []byte
data, err = os.ReadFile(filePath) data, err = os.ReadFile(filePath)
if err != nil { if err != nil {
return nil, fmt.Errorf("error reading file %s: %w", filePath, err) return fmt.Errorf("error reading file %s: %w", filePath, err)
} }
// Parse JSON // Parse JSON
var provider Provider var provider Provider
if provider, err = decodeProviderConfig(data); err != nil { if provider, err = decodeProviderConfig(data); err != nil {
return nil, fmt.Errorf("error parsing JSON from file %s: %w", filePath, err) return fmt.Errorf("error parsing JSON from file %s: %w", filePath, err)
} }
for _, model := range provider.Models { for _, model := range provider.Models {
@@ -275,7 +281,7 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) {
provider.ModelDriver, err = modelFactory.CreateModelDriver(provider.Name, provider.URL, provider.URLSuffix) provider.ModelDriver, err = modelFactory.CreateModelDriver(provider.Name, provider.URL, provider.URLSuffix)
if err != nil { if err != nil {
return nil, fmt.Errorf("error creating model driver for provider %s: %w", provider.Name, err) return fmt.Errorf("error creating model driver for provider %s: %w", provider.Name, err)
} }
// Add to providers list // Add to providers list
@@ -283,14 +289,14 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) {
} }
if len(providers) == 0 { if len(providers) == 0 {
return nil, fmt.Errorf("no JSON files found in directory %s", dirPath) return fmt.Errorf("no JSON files found in directory %s", dirPath)
} }
// Read the file // Read the file
var data []byte var data []byte
data, err = os.ReadFile("conf/all_models.json") data, err = os.ReadFile("conf/all_models.json")
if err != nil { if err != nil {
return nil, fmt.Errorf("error reading file 'conf/all_models.json': %w", err) return fmt.Errorf("error reading file 'conf/all_models.json': %w", err)
} }
// Parse JSON // Parse JSON
@@ -299,7 +305,7 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) {
} }
var allModels AllModels var allModels AllModels
if err = json.Unmarshal(data, &allModels); err != nil { if err = json.Unmarshal(data, &allModels); err != nil {
return nil, fmt.Errorf("error parsing JSON from file 'conf/all_models.json': %w", err) return fmt.Errorf("error parsing JSON from file 'conf/all_models.json': %w", err)
} }
alias2ModelIndex := make(map[string]int) alias2ModelIndex := make(map[string]int)
@@ -313,11 +319,12 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) {
} }
} }
return &ProviderManager{ providerManager = &ProviderManager{
Providers: providers, Providers: providers,
AllModels: allModels.Models, AllModels: allModels.Models,
Alias2ModelIndex: alias2ModelIndex, Alias2ModelIndex: alias2ModelIndex,
}, nil }
return nil
} }
// 1. List all providers // 1. List all providers
@@ -371,8 +378,8 @@ func (pm *ProviderManager) ListAllModels() ([]map[string]interface{}, error) {
if model.Thinking != nil { if model.Thinking != nil {
modelData["thinking"] = model.Thinking modelData["thinking"] = model.Thinking
} }
if model.MaxTokens != 0 { if model.MaxTokens != nil {
modelData["max_tokens"] = model.MaxTokens modelData["max_tokens"] = *model.MaxTokens
} }
modelList = append(modelList, modelData) modelList = append(modelList, modelData)
} }
@@ -385,9 +392,13 @@ func (pm *ProviderManager) ListAllModels() ([]map[string]interface{}, error) {
} }
func (pm *ProviderManager) GetModelByNameOrAlias(modelName string) *Model { func (pm *ProviderManager) GetModelByNameOrAlias(modelName string) *Model {
lowerModelName := strings.ToLower(modelName)
// Check if it is alias // Check if it is alias
modelIndex := pm.Alias2ModelIndex[modelName] modelIndex, ok := pm.Alias2ModelIndex[lowerModelName]
return &pm.AllModels[modelIndex] if ok {
return &pm.AllModels[modelIndex]
}
return nil
} }
// 2. Show specific provider information (including base_url) // 2. Show specific provider information (including base_url)
@@ -504,7 +515,7 @@ func (pm *ProviderManager) SearchModelInfo(providerName, modelName string, filte
switch filterBy { switch filterBy {
case "max_tokens": case "max_tokens":
if maxVal, ok := filterValue.(int); ok { if maxVal, ok := filterValue.(int); ok {
if model.MaxTokens < maxVal { if *model.MaxTokens < maxVal {
matchFilter = false matchFilter = false
resp.Code = 400 resp.Code = 400
resp.Message = fmt.Sprintf("Model does not meet filter criteria: max_tokens (%d) < %d", resp.Message = fmt.Sprintf("Model does not meet filter criteria: max_tokens (%d) < %d",
@@ -14,12 +14,11 @@
// limitations under the License. // limitations under the License.
// //
package entity package models
import ( import (
"os" "os"
"path/filepath" "path/filepath"
modeldrivers "ragflow/internal/entity/models"
"strings" "strings"
"testing" "testing"
) )
@@ -54,16 +53,18 @@ func TestHostedProviderConfigsLoadSharedDrivers(t *testing.T) {
} }
} }
pm, err := NewProviderManager(dir) err := InitProviderManager(dir)
if err != nil { if err != nil {
t.Fatalf("NewProviderManager: %v", err) t.Fatalf("InitProviderManager: %v", err)
} }
pm := GetProviderManager()
minerU := pm.FindProvider("MinerU.Net") minerU := pm.FindProvider("MinerU.Net")
if minerU == nil { if minerU == nil {
t.Fatal("MinerU.Net provider not found") t.Fatal("MinerU.Net provider not found")
} }
if _, ok := minerU.ModelDriver.(*modeldrivers.MinerUModel); !ok { if _, ok := minerU.ModelDriver.(*MinerUModel); !ok {
t.Fatalf("MinerU.Net ModelDriver=%T, want *models.MinerUModel", minerU.ModelDriver) t.Fatalf("MinerU.Net ModelDriver=%T, want *models.MinerUModel", minerU.ModelDriver)
} }
if minerU.Class != "mineru.net" { if minerU.Class != "mineru.net" {
@@ -77,7 +78,7 @@ func TestHostedProviderConfigsLoadSharedDrivers(t *testing.T) {
if paddleOCR == nil { if paddleOCR == nil {
t.Fatal("PaddleOCR.Net provider not found") t.Fatal("PaddleOCR.Net provider not found")
} }
if _, ok := paddleOCR.ModelDriver.(*modeldrivers.PaddleOCRModel); !ok { if _, ok := paddleOCR.ModelDriver.(*PaddleOCRModel); !ok {
t.Fatalf("PaddleOCR.Net ModelDriver=%T, want *models.PaddleOCRModel", paddleOCR.ModelDriver) t.Fatalf("PaddleOCR.Net ModelDriver=%T, want *models.PaddleOCRModel", paddleOCR.ModelDriver)
} }
if paddleOCR.Class != "paddleocr.net" { if paddleOCR.Class != "paddleocr.net" {
@@ -96,16 +97,18 @@ func TestLocalOCRProviderConfigsLoadLocalDrivers(t *testing.T) {
} }
} }
pm, err := NewProviderManager(dir) err := InitProviderManager(dir)
if err != nil { if err != nil {
t.Fatalf("NewProviderManager: %v", err) t.Fatalf("InitProviderManager: %v", err)
} }
pm := GetProviderManager()
minerU := pm.FindProvider("MinerU") minerU := pm.FindProvider("MinerU")
if minerU == nil { if minerU == nil {
t.Fatal("MinerU provider not found") t.Fatal("MinerU provider not found")
} }
if _, ok := minerU.ModelDriver.(*modeldrivers.MinerULocalModel); !ok { if _, ok := minerU.ModelDriver.(*MinerULocalModel); !ok {
t.Fatalf("MinerU ModelDriver=%T, want *models.MinerULocalModel", minerU.ModelDriver) t.Fatalf("MinerU ModelDriver=%T, want *models.MinerULocalModel", minerU.ModelDriver)
} }
if minerU.URLSuffix.DocumentParse != "file_parse" { if minerU.URLSuffix.DocumentParse != "file_parse" {
@@ -116,7 +119,7 @@ func TestLocalOCRProviderConfigsLoadLocalDrivers(t *testing.T) {
if paddleOCR == nil { if paddleOCR == nil {
t.Fatal("PaddleOCR provider not found") t.Fatal("PaddleOCR provider not found")
} }
if _, ok := paddleOCR.ModelDriver.(*modeldrivers.PaddleOCRLocalModel); !ok { if _, ok := paddleOCR.ModelDriver.(*PaddleOCRLocalModel); !ok {
t.Fatalf("PaddleOCR ModelDriver=%T, want *models.PaddleOCRLocalModel", paddleOCR.ModelDriver) t.Fatalf("PaddleOCR ModelDriver=%T, want *models.PaddleOCRLocalModel", paddleOCR.ModelDriver)
} }
if paddleOCR.URLSuffix.OCR != "layout-parsing" { if paddleOCR.URLSuffix.OCR != "layout-parsing" {
@@ -132,9 +135,9 @@ func TestProviderConfigsLoadURLSuffixKeys(t *testing.T) {
} }
} }
pm, err := NewProviderManager(dir) err := InitProviderManager(dir)
if err != nil { if err != nil {
t.Fatalf("NewProviderManager: %v", err) t.Fatalf("InitProviderManager: %v", err)
} }
cohere := pm.FindProvider("CoHere") cohere := pm.FindProvider("CoHere")
@@ -177,9 +180,9 @@ func TestProviderConfigRejectsUnknownURLSuffixKey(t *testing.T) {
t.Fatalf("write config: %v", err) t.Fatalf("write config: %v", err)
} }
_, err := NewProviderManager(dir) _, err := InitProviderManager(dir)
if err == nil { if err == nil {
t.Fatal("NewProviderManager succeeded with unknown url_suffix key") t.Fatal("InitProviderManager succeeded with unknown url_suffix key")
} }
if !strings.Contains(err.Error(), `unknown field "unknown_suffix"`) { if !strings.Contains(err.Error(), `unknown field "unknown_suffix"`) {
t.Fatalf("error=%q, want unknown_suffix field", err) t.Fatalf("error=%q, want unknown_suffix field", err)
@@ -195,9 +198,9 @@ func TestPPIOProviderConfigLoadsIntoProviderManager(t *testing.T) {
t.Fatalf("write ppio config: %v", err) t.Fatalf("write ppio config: %v", err)
} }
pm, err := NewProviderManager(dir) err := InitProviderManager(dir)
if err != nil { if err != nil {
t.Fatalf("NewProviderManager: %v", err) t.Fatalf("InitProviderManager: %v", err)
} }
provider := pm.FindProvider("ppio") provider := pm.FindProvider("ppio")
@@ -219,7 +222,7 @@ func TestPPIOProviderConfigLoadsIntoProviderManager(t *testing.T) {
if provider.URLSuffix.Models != "models" { if provider.URLSuffix.Models != "models" {
t.Errorf("models suffix=%q", provider.URLSuffix.Models) t.Errorf("models suffix=%q", provider.URLSuffix.Models)
} }
if _, ok := provider.ModelDriver.(*modeldrivers.PPIOModel); !ok { if _, ok := provider.ModelDriver.(*PPIOModel); !ok {
t.Fatalf("ModelDriver=%T, want *models.PPIOModel", provider.ModelDriver) t.Fatalf("ModelDriver=%T, want *models.PPIOModel", provider.ModelDriver)
} }
if provider.ModelDriver.Name() != "ppio" { if provider.ModelDriver.Name() != "ppio" {
@@ -282,9 +285,9 @@ func TestSiliconFlowProviderConfigLoadsLatestProModels(t *testing.T) {
t.Fatalf("write siliconflow config: %v", err) t.Fatalf("write siliconflow config: %v", err)
} }
pm, err := NewProviderManager(dir) err := InitProviderManager(dir)
if err != nil { if err != nil {
t.Fatalf("NewProviderManager: %v", err) t.Fatalf("InitProviderManager: %v", err)
} }
provider := pm.FindProvider("SiliconFlow") provider := pm.FindProvider("SiliconFlow")
@@ -297,7 +300,7 @@ func TestSiliconFlowProviderConfigLoadsLatestProModels(t *testing.T) {
if provider.URLSuffix.Chat != "chat/completions" { if provider.URLSuffix.Chat != "chat/completions" {
t.Errorf("chat suffix=%q", provider.URLSuffix.Chat) t.Errorf("chat suffix=%q", provider.URLSuffix.Chat)
} }
if _, ok := provider.ModelDriver.(*modeldrivers.SiliconflowModel); !ok { if _, ok := provider.ModelDriver.(*SiliconflowModel); !ok {
t.Fatalf("ModelDriver=%T, want *models.SiliconflowModel", provider.ModelDriver) t.Fatalf("ModelDriver=%T, want *models.SiliconflowModel", provider.ModelDriver)
} }
if provider.ModelDriver.Name() != "siliconflow" { if provider.ModelDriver.Name() != "siliconflow" {
+3 -3
View File
@@ -402,7 +402,7 @@ func (m *ModelScopeModel) ParseFile(modelName *string, content []byte, url *stri
// ListModels returns the model IDs exposed by ModelScope's OpenAI-compatible // ListModels returns the model IDs exposed by ModelScope's OpenAI-compatible
// /v1/models endpoint. // /v1/models endpoint.
func (m *ModelScopeModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (m *ModelScopeModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := m.baseModel.APIConfigCheck(apiConfig); err != nil { if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -446,10 +446,10 @@ func (m *ModelScopeModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{Name: model.ID})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -385,7 +385,7 @@ func (m *MoonshotModel) Embed(modelName *string, texts []string, apiConfig *APIC
return nil, fmt.Errorf("not implemented") return nil, fmt.Errorf("not implemented")
} }
func (m *MoonshotModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (m *MoonshotModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := m.baseModel.APIConfigCheck(apiConfig); err != nil { if err := m.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -435,12 +435,12 @@ func (m *MoonshotModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("models response missing data") return nil, fmt.Errorf("models response missing data")
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if strings.TrimSpace(model.ID) == "" { if strings.TrimSpace(model.ID) == "" {
return nil, fmt.Errorf("models response contains empty id") return nil, fmt.Errorf("models response contains empty id")
} }
models = append(models, strings.TrimSpace(model.ID)) models = append(models, ListModelResponse{Name: strings.TrimSpace(model.ID)})
} }
return models, nil return models, nil
+3 -3
View File
@@ -532,7 +532,7 @@ type n1nModelCatalogResponse struct {
} }
// ListModels returns the live n1n.ai model catalog // ListModels returns the live n1n.ai model catalog
func (n *N1NModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (n *N1NModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := n.baseModel.APIConfigCheck(apiConfig); err != nil { if err := n.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -570,10 +570,10 @@ func (n *N1NModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(parsed.Data)) models := make([]ListModelResponse, 0, len(parsed.Data))
for _, item := range parsed.Data { for _, item := range parsed.Data {
if item.ID != "" { if item.ID != "" {
models = append(models, item.ID) models = append(models, ListModelResponse{Name: item.ID})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -526,7 +526,7 @@ func (n *NovitaModel) ChatStreamlyWithSender(modelName string, messages []Messag
} }
// ListModels returns the list of model ids visible to the API key. // ListModels returns the list of model ids visible to the API key.
func (n *NovitaModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (n *NovitaModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := n.baseModel.APIConfigCheck(apiConfig); err != nil { if err := n.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -572,7 +572,7 @@ func (n *NovitaModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0) models := make([]ListModelResponse, 0, len(data))
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -582,7 +582,7 @@ func (n *NovitaModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
} }
+3 -3
View File
@@ -591,7 +591,7 @@ func (n *NvidiaModel) ParseFile(modelName *string, content []byte, url *string,
// and returns the list of available model ids. The endpoint is // and returns the list of available model ids. The endpoint is
// OpenAI-compatible, so the parsing follows the same shape used by // OpenAI-compatible, so the parsing follows the same shape used by
// the moonshot, xai, and openai drivers. // the moonshot, xai, and openai drivers.
func (n NvidiaModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (n NvidiaModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := n.baseModel.APIConfigCheck(apiConfig); err != nil { if err := n.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -642,7 +642,7 @@ func (n NvidiaModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0, len(data)) models := make([]ListModelResponse, 0, len(data))
for _, item := range data { for _, item := range data {
m, ok := item.(map[string]interface{}) m, ok := item.(map[string]interface{})
if !ok { if !ok {
@@ -652,7 +652,7 @@ func (n NvidiaModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, id) models = append(models, ListModelResponse{Name: id})
} }
return models, nil return models, nil
+3 -3
View File
@@ -461,7 +461,7 @@ func (o *OllamaModel) ParseFile(modelName *string, content []byte, url *string,
return nil, fmt.Errorf("%s, no such method", o.Name()) return nil, fmt.Errorf("%s, no such method", o.Name())
} }
func (o *OllamaModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (o *OllamaModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
resolvedBaseURL, err := o.baseModel.GetBaseURL(apiConfig) resolvedBaseURL, err := o.baseModel.GetBaseURL(apiConfig)
if err != nil { if err != nil {
@@ -515,11 +515,11 @@ func (o *OllamaModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
// convert result["data"] to []map[string]interface{} // convert result["data"] to []map[string]interface{}
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["models"].([]interface{}) { for _, model := range result["models"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["name"].(string) modelName := modelMap["name"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -444,7 +444,7 @@ func (o *OpenAIModel) Embed(modelName *string, texts []string, apiConfig *APICon
} }
// ListModels returns the list of model ids visible to the API key. // ListModels returns the list of model ids visible to the API key.
func (o *OpenAIModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (o *OpenAIModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := o.baseModel.APIConfigCheck(apiConfig); err != nil { if err := o.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -493,7 +493,7 @@ func (o *OpenAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -503,7 +503,7 @@ func (o *OpenAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -740,7 +740,7 @@ func (o *OpenRouterModel) ParseFile(modelName *string, content []byte, url *stri
return nil, fmt.Errorf("%s, no such method", o.Name()) return nil, fmt.Errorf("%s, no such method", o.Name())
} }
func (o *OpenRouterModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (o *OpenRouterModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := o.baseModel.APIConfigCheck(apiConfig); err != nil { if err := o.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -792,11 +792,11 @@ func (o *OpenRouterModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
// convert result["data"] to []map[string]interface{} // convert result["data"] to []map[string]interface{}
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["data"].([]interface{}) { for _, model := range result["data"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["id"].(string) modelName := modelMap["id"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -413,7 +413,7 @@ func (o *OrcaRouterModel) ParseFile(modelName *string, content []byte, url *stri
return nil, fmt.Errorf("%s no such method", o.Name()) return nil, fmt.Errorf("%s no such method", o.Name())
} }
func (o *OrcaRouterModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (o *OrcaRouterModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := o.baseModel.APIConfigCheck(apiConfig); err != nil { if err := o.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -464,11 +464,11 @@ func (o *OrcaRouterModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
// convert result["data"] to []map[string]interface{} // convert result["data"] to []map[string]interface{}
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["data"].([]interface{}) { for _, model := range result["data"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["id"].(string) modelName := modelMap["id"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+4 -4
View File
@@ -212,9 +212,9 @@ func (p *PaddleOCRModel) OCRFile(modelName *string, content []byte, fileURL *str
} }
pollReq, _ := http.NewRequestWithContext(ctx, "GET", pollUrl, nil) pollReq, _ := http.NewRequestWithContext(ctx, "GET", pollUrl, nil)
if auth := BearerAuth(apiConfig); auth != "" { if auth := BearerAuth(apiConfig); auth != "" {
pollReq.Header.Set("Authorization", auth) pollReq.Header.Set("Authorization", auth)
} }
pollResp, err := p.baseModel.httpClient.Do(pollReq) pollResp, err := p.baseModel.httpClient.Do(pollReq)
if err != nil { if err != nil {
@@ -296,7 +296,7 @@ func (p *PaddleOCRModel) ParseFile(modelName *string, content []byte, url *strin
return nil, fmt.Errorf("%s, no such method", p.Name()) return nil, fmt.Errorf("%s, no such method", p.Name())
} }
func (p *PaddleOCRModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (p *PaddleOCRModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("%s, no such method", p.Name()) return nil, fmt.Errorf("%s, no such method", p.Name())
} }
+1 -1
View File
@@ -190,7 +190,7 @@ func (p *PaddleOCRLocalModel) ParseFile(modelName *string, content []byte, url *
return nil, fmt.Errorf("%s no such method", p.Name()) return nil, fmt.Errorf("%s no such method", p.Name())
} }
func (p *PaddleOCRLocalModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (p *PaddleOCRLocalModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("%s no such method", p.Name()) return nil, fmt.Errorf("%s no such method", p.Name())
} }
+5 -5
View File
@@ -311,7 +311,7 @@ type perplexityModelListResponse struct {
Data []perplexityModelInfo `json:"data"` Data []perplexityModelInfo `json:"data"`
} }
func (p *PerplexityModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (p *PerplexityModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := p.baseModel.APIConfigCheck(apiConfig); err != nil { if err := p.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -350,10 +350,10 @@ func (p *PerplexityModel) ListModels(apiConfig *APIConfig) ([]string, error) {
// or alternate payloads may return a bare array; accept both. // or alternate payloads may return a bare array; accept both.
var wrapped perplexityModelListResponse var wrapped perplexityModelListResponse
if err = json.Unmarshal(body, &wrapped); err == nil && len(wrapped.Data) > 0 { if err = json.Unmarshal(body, &wrapped); err == nil && len(wrapped.Data) > 0 {
models := make([]string, 0, len(wrapped.Data)) models := make([]ListModelResponse, 0, len(wrapped.Data))
for _, model := range wrapped.Data { for _, model := range wrapped.Data {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{Name: model.ID})
} }
} }
return models, nil return models, nil
@@ -363,10 +363,10 @@ func (p *PerplexityModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if err = json.Unmarshal(body, &bare); err != nil { if err = json.Unmarshal(body, &bare); err != nil {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(bare)) models := make([]ListModelResponse, 0, len(bare))
for _, model := range bare { for _, model := range bare {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{Name: model.ID})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -314,7 +314,7 @@ type ppioListModelsResponse struct {
Error interface{} `json:"error"` Error interface{} `json:"error"`
} }
func (p *PPIOModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (p *PPIOModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := p.baseModel.APIConfigCheck(apiConfig); err != nil { if err := p.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -356,10 +356,10 @@ func (p *PPIOModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("ppio: upstream error: %v", result.Error) return nil, fmt.Errorf("ppio: upstream error: %v", result.Error)
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{Name: model.ID})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -445,7 +445,7 @@ func (q *QiniuModel) ParseFile(modelName *string, content []byte, url *string, a
return nil, fmt.Errorf("%s, no such method", q.Name()) return nil, fmt.Errorf("%s, no such method", q.Name())
} }
func (q *QiniuModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (q *QiniuModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := q.baseModel.APIConfigCheck(apiConfig); err != nil { if err := q.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -490,7 +490,7 @@ func (q *QiniuModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0, len(data)) models := make([]ListModelResponse, 0, len(data))
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -500,7 +500,7 @@ func (q *QiniuModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok || modelID == "" { if !ok || modelID == "" {
continue continue
} }
models = append(models, modelID) models = append(models, ListModelResponse{Name: modelID})
} }
return models, nil return models, nil
+3 -3
View File
@@ -504,7 +504,7 @@ func dispatchReplicateSSEEvent(event replicateSSEEvent, sender func(*string, *st
} }
} }
func (r *ReplicateModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (r *ReplicateModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := r.baseModel.APIConfigCheck(apiConfig); err != nil { if err := r.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -543,10 +543,10 @@ func (r *ReplicateModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(result.Results)) models := make([]ListModelResponse, 0, len(result.Results))
for _, model := range result.Results { for _, model := range result.Results {
if model.Owner != "" && model.Name != "" { if model.Owner != "" && model.Name != "" {
models = append(models, fmt.Sprintf("%s/%s", model.Owner, model.Name)) models = append(models, ListModelResponse{Name: model.Owner + "/" + model.Name})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -495,7 +495,7 @@ func (s *SiliconflowModel) Embed(modelName *string, texts []string, apiConfig *A
return embeddings, nil return embeddings, nil
} }
func (s *SiliconflowModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (s *SiliconflowModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := s.baseModel.APIConfigCheck(apiConfig); err != nil { if err := s.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -546,13 +546,13 @@ func (s *SiliconflowModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
var models []string var models []ListModelResponse
for _, model := range modelList.Models { for _, model := range modelList.Models {
modelName := model.ID modelName := model.ID
if model.OwnedBy != "" { if model.OwnedBy != "" {
modelName = model.ID + "@" + model.OwnedBy modelName = model.ID + "@" + model.OwnedBy
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -321,7 +321,7 @@ func (s *StepFunModel) Embed(modelName *string, texts []string, apiConfig *APICo
} }
// ListModels returns the list of model ids visible to the API key. // ListModels returns the list of model ids visible to the API key.
func (s *StepFunModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (s *StepFunModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := s.baseModel.APIConfigCheck(apiConfig); err != nil { if err := s.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -368,7 +368,7 @@ func (s *StepFunModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -378,7 +378,7 @@ func (s *StepFunModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -316,7 +316,7 @@ type togetherAIModelInfo struct {
ID string `json:"id"` ID string `json:"id"`
} }
func (t *TogetherAIModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (t *TogetherAIModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := t.baseModel.APIConfigCheck(apiConfig); err != nil { if err := t.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -357,10 +357,10 @@ func (t *TogetherAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(result)) models := make([]ListModelResponse, 0, len(result))
for _, model := range result { for _, model := range result {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{Name: model.ID})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -460,7 +460,7 @@ func (t *TokenHubModel) ParseFile(modelName *string, content []byte, url *string
return nil, fmt.Errorf("%s no such method", t.Name()) return nil, fmt.Errorf("%s no such method", t.Name())
} }
func (t *TokenHubModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (t *TokenHubModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := t.baseModel.APIConfigCheck(apiConfig); err != nil { if err := t.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -507,7 +507,7 @@ func (t *TokenHubModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0, len(data)) models := make([]ListModelResponse, 0, len(data))
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -517,7 +517,7 @@ func (t *TokenHubModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -312,7 +312,7 @@ func (t *TokenPonyModel) ChatStreamlyWithSender(modelName string, messages []Mes
return nil return nil
} }
func (t *TokenPonyModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (t *TokenPonyModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := t.baseModel.APIConfigCheck(apiConfig); err != nil { if err := t.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -357,7 +357,7 @@ func (t *TokenPonyModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0, len(data)) models := make([]ListModelResponse, 0, len(data))
for _, m := range data { for _, m := range data {
modelMap, ok := m.(map[string]interface{}) modelMap, ok := m.(map[string]interface{})
if !ok { if !ok {
@@ -367,7 +367,7 @@ func (t *TokenPonyModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, id) models = append(models, ListModelResponse{Name: id})
} }
return models, nil return models, nil
} }
+9 -1
View File
@@ -36,7 +36,7 @@ type ModelDriver interface {
// ParseFile parse file // ParseFile parse file
ParseFile(modelName *string, content []byte, url *string, apiConfig *APIConfig, parseFileConfig *ParseFileConfig) (*ParseFileResponse, error) ParseFile(modelName *string, content []byte, url *string, apiConfig *APIConfig, parseFileConfig *ParseFileConfig) (*ParseFileResponse, error)
// ListModels List supported models // ListModels List supported models
ListModels(apiConfig *APIConfig) ([]string, error) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error)
Balance(apiConfig *APIConfig) (map[string]interface{}, error) Balance(apiConfig *APIConfig) (map[string]interface{}, error)
@@ -78,6 +78,14 @@ type OCRFileResponse struct {
Text *string `json:"text"` Text *string `json:"text"`
} }
type ListModelResponse struct {
Name string `json:"name"`
MaxTokens *int `json:"max_tokens"`
ModelTypes []string `json:"model_types"`
Thinking *ModelThinking `json:"thinking"`
Dimension *int `json:"dimension"` // used by embedding models
}
type ParseFileResponse struct { type ParseFileResponse struct {
TaskID string `json:"task_id"` TaskID string `json:"task_id"`
} }
+3 -3
View File
@@ -429,7 +429,7 @@ func (u *UpstageModel) Embed(modelName *string, texts []string, apiConfig *APICo
} }
// ListModels returns the list of model ids visible to the API key. // ListModels returns the list of model ids visible to the API key.
func (u *UpstageModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (u *UpstageModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := u.baseModel.APIConfigCheck(apiConfig); err != nil { if err := u.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -476,7 +476,7 @@ func (u *UpstageModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -486,7 +486,7 @@ func (u *UpstageModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -476,7 +476,7 @@ func (v *VllmModel) Embed(modelName *string, texts []string, apiConfig *APIConfi
return embeddings, nil return embeddings, nil
} }
func (v *VllmModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (v *VllmModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := v.baseModel.APIConfigCheck(apiConfig); err != nil { if err := v.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -537,11 +537,11 @@ func (v *VllmModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
// convert result["data"] to []map[string]interface{} // convert result["data"] to []map[string]interface{}
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["data"].([]interface{}) { for _, model := range result["data"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["id"].(string) modelName := modelMap["id"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -563,7 +563,7 @@ func (v *VolcEngine) ParseFile(modelName *string, content []byte, url *string, a
return nil, fmt.Errorf("%s, no such method", v.Name()) return nil, fmt.Errorf("%s, no such method", v.Name())
} }
func (v *VolcEngine) ListModels(apiConfig *APIConfig) ([]string, error) { func (v *VolcEngine) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := v.baseModel.APIConfigCheck(apiConfig); err != nil { if err := v.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -613,13 +613,13 @@ func (v *VolcEngine) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(modelList.Models)) models := make([]ListModelResponse, 0, len(modelList.Models))
for _, model := range modelList.Models { for _, model := range modelList.Models {
modelName := model.ID modelName := model.ID
if model.OwnedBy != "" { if model.OwnedBy != "" {
modelName = model.ID + "@" + model.OwnedBy modelName = model.ID + "@" + model.OwnedBy
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+1 -1
View File
@@ -270,7 +270,7 @@ func (v *VoyageModel) Rerank(modelName *string, query string, documents []string
return rerankResponse, nil return rerankResponse, nil
} }
func (v *VoyageModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (v *VoyageModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("%s, no such method", v.Name()) return nil, fmt.Errorf("%s, no such method", v.Name())
} }
+3 -3
View File
@@ -364,7 +364,7 @@ func (x *XAIModel) Embed(modelName *string, texts []string, apiConfig *APIConfig
} }
// ListModels returns the list of model ids visible to the API key. // ListModels returns the list of model ids visible to the API key.
func (x *XAIModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (x *XAIModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := x.baseModel.APIConfigCheck(apiConfig); err != nil { if err := x.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -417,7 +417,7 @@ func (x *XAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("invalid models list format") return nil, fmt.Errorf("invalid models list format")
} }
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range data { for _, model := range data {
modelMap, ok := model.(map[string]interface{}) modelMap, ok := model.(map[string]interface{})
if !ok { if !ok {
@@ -427,7 +427,7 @@ func (x *XAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
if !ok { if !ok {
continue continue
} }
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+1 -1
View File
@@ -884,7 +884,7 @@ func (x *XiaomiModel) ParseFile(modelName *string, content []byte, url *string,
return nil, fmt.Errorf("no such method %s", x.Name()) return nil, fmt.Errorf("no such method %s", x.Name())
} }
func (x *XiaomiModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (x *XiaomiModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
return nil, fmt.Errorf("no such method %s", x.Name()) return nil, fmt.Errorf("no such method %s", x.Name())
} }
+3 -3
View File
@@ -757,7 +757,7 @@ func (x *XinferenceModel) ParseFile(modelName *string, content []byte, url *stri
// ListModels returns the model IDs exposed by Xinference's OpenAI-compatible // ListModels returns the model IDs exposed by Xinference's OpenAI-compatible
// /v1/models endpoint. // /v1/models endpoint.
func (x *XinferenceModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (x *XinferenceModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := x.baseModel.APIConfigCheck(apiConfig); err != nil { if err := x.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -801,10 +801,10 @@ func (x *XinferenceModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(result.Data)) models := make([]ListModelResponse, 0, len(result.Data))
for _, model := range result.Data { for _, model := range result.Data {
if model.ID != "" { if model.ID != "" {
models = append(models, model.ID) models = append(models, ListModelResponse{Name: model.ID})
} }
} }
return models, nil return models, nil
+3 -3
View File
@@ -384,7 +384,7 @@ func (x *XunFeiModel) ParseFile(modelName *string, content []byte, url *string,
return nil, fmt.Errorf("%s, no such method", x.Name()) return nil, fmt.Errorf("%s, no such method", x.Name())
} }
func (x *XunFeiModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (x *XunFeiModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := x.baseModel.APIConfigCheck(apiConfig); err != nil { if err := x.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -436,11 +436,11 @@ func (x *XunFeiModel) ListModels(apiConfig *APIConfig) ([]string, error) {
} }
// convert result["data"] to []map[string]interface{} // convert result["data"] to []map[string]interface{}
models := make([]string, 0) models := make([]ListModelResponse, 0)
for _, model := range result["data"].([]interface{}) { for _, model := range result["data"].([]interface{}) {
modelMap := model.(map[string]interface{}) modelMap := model.(map[string]interface{})
modelName := modelMap["id"].(string) modelName := modelMap["id"].(string)
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+3 -3
View File
@@ -474,7 +474,7 @@ func (z *ZhipuAIModel) Embed(modelName *string, texts []string, apiConfig *APICo
return embeddings, nil return embeddings, nil
} }
func (z *ZhipuAIModel) ListModels(apiConfig *APIConfig) ([]string, error) { func (z *ZhipuAIModel) ListModels(apiConfig *APIConfig) ([]ListModelResponse, error) {
if err := z.baseModel.APIConfigCheck(apiConfig); err != nil { if err := z.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err return nil, err
} }
@@ -517,10 +517,10 @@ func (z *ZhipuAIModel) ListModels(apiConfig *APIConfig) ([]string, error) {
return nil, fmt.Errorf("failed to parse response: %w", err) return nil, fmt.Errorf("failed to parse response: %w", err)
} }
models := make([]string, 0, len(modelList.Models)) models := make([]ListModelResponse, 0, len(modelList.Models))
for _, model := range modelList.Models { for _, model := range modelList.Models {
modelName := model.ID modelName := model.ID
models = append(models, modelName) models = append(models, ListModelResponse{Name: modelName})
} }
return models, nil return models, nil
+1 -8
View File
@@ -680,17 +680,10 @@ func (h *ProviderHandler) ListInstanceModels(c *gin.Context) {
return return
} }
var modelResponse []map[string]string
for _, modelName := range modelList {
modelResponse = append(modelResponse, map[string]string{
"model_name": modelName,
})
}
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"code": 0, "code": 0,
"message": "success", "message": "success",
"data": modelResponse, "data": modelList,
}) })
return return
} }
+50 -19
View File
@@ -184,7 +184,7 @@ func (m *ModelProviderService) DeleteModelProvider(providerName, userID string)
return common.CodeSuccess, nil 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 // Get tenant ID from user
tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") 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) { 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") return common.CodeServerError, errors.New("fail to get UUID")
} }
var modelSchema *entity.Model var modelSchema *modelModule.Model
modelSchema, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName) modelSchema, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) 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") 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) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) 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 return common.CodeNotFound, err
} }
var model *entity.Model = nil var model *modelModule.Model = nil
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return common.CodeNotFound, err return common.CodeNotFound, err
@@ -1132,7 +1147,7 @@ func (m *ModelProviderService) EmbedText(providerName, instanceName, modelName,
return nil, common.CodeNotFound, errors.New("provider not found") 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) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) 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") 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) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) 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") 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) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) 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 return common.CodeNotFound, err
} }
var model *entity.Model = nil var model *modelModule.Model = nil
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return common.CodeNotFound, err return common.CodeNotFound, err
@@ -1553,7 +1568,7 @@ func (m *ModelProviderService) AudioSpeech(providerName, instanceName, modelName
return nil, common.CodeNotFound, errors.New("provider not found") 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) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) 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 return common.CodeNotFound, err
} }
var model *entity.Model = nil var model *modelModule.Model = nil
model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return common.CodeNotFound, err 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") 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) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) 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") 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) model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName)
if err != nil { if err != nil {
return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) 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: &region} apiConfig := &modelModule.APIConfig{ApiKey: &apiKey, Region: &region}
maxTokens := 0 maxTokens := 0
if mi, _ := dao.GetModelProviderManager().GetModelByName("Builtin", pureModelName); mi != nil { 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 return builtinDriver, pureModelName, apiConfig, maxTokens, nil
} }
@@ -2275,7 +2294,11 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin
} }
maxTokens := 0 maxTokens := 0
if mi, _ := dao.GetModelProviderManager().GetModelByName(providerName, pureModelName); mi != nil { 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: &region, BaseURL: &baseURL} apiConfig := &modelModule.APIConfig{ApiKey: &apiKey, Region: &region, BaseURL: &baseURL}
return driver, modelObj.ModelName, apiConfig, maxTokens, nil return driver, modelObj.ModelName, apiConfig, maxTokens, nil
@@ -2309,7 +2332,7 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin
if targetProvider == nil { if targetProvider == nil {
return nil, "", nil, 0, fmt.Errorf("model provider config not found: %s", providerName) 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 { for i := range targetProvider.Models {
if strings.EqualFold(targetProvider.Models[i].Name, pureModelName) { if strings.EqualFold(targetProvider.Models[i].Name, pureModelName) {
llmInfo = targetProvider.Models[i] llmInfo = targetProvider.Models[i]
@@ -2324,7 +2347,11 @@ func (m *ModelProviderService) GetModelConfigFromProviderInstance(tenantID strin
return nil, "", nil, 0, driverErr return nil, "", nil, 0, driverErr
} }
apiConfig := &modelModule.APIConfig{ApiKey: &apiKey, Region: &region, BaseURL: &baseURL} apiConfig := &modelModule.APIConfig{ApiKey: &apiKey, Region: &region, 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 // 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) modelInfo, err := dao.GetModelProviderManager().GetModelByName(providerName, modelName)
maxTokens := 0 maxTokens := 0
if err == nil && modelInfo != nil { 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 // 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 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 return dao.GetModelProviderManager().GetModelByNameOrAlias(modelName), nil
} }