fix(playground): resolve auto group model listing (#6163)
* fix(playground): resolve auto group model listing - merge and deduplicate available models in configured auto group order. - reuse special usable group rules and add model filtering regression coverage. * refactor: extract GetGroupsEnabledModels to dedupe group model expansion
This commit is contained in:
+1
-13
@@ -245,19 +245,7 @@ func ListModels(c *gin.Context, modelType int) {
|
||||
userModelNames = append(userModelNames, allowModel)
|
||||
}
|
||||
} else {
|
||||
var models []string
|
||||
if groups.tokenGroup == "auto" {
|
||||
for _, autoGroup := range ownerGroups {
|
||||
groupModels := model.GetGroupEnabledModels(autoGroup)
|
||||
for _, g := range groupModels {
|
||||
if !common.StringsContains(models, g) {
|
||||
models = append(models, g)
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
models = model.GetGroupEnabledModels(ownerGroups[0])
|
||||
}
|
||||
models := service.GetGroupsEnabledModels(ownerGroups)
|
||||
for _, modelName := range models {
|
||||
if !acceptUnsetRatioModel {
|
||||
if !helper.HasModelBillingConfig(modelName) {
|
||||
|
||||
@@ -12,8 +12,10 @@ import (
|
||||
"github.com/QuantumNous/new-api/constant"
|
||||
"github.com/QuantumNous/new-api/dto"
|
||||
"github.com/QuantumNous/new-api/model"
|
||||
"github.com/QuantumNous/new-api/setting"
|
||||
"github.com/QuantumNous/new-api/setting/config"
|
||||
"github.com/QuantumNous/new-api/setting/operation_setting"
|
||||
"github.com/QuantumNous/new-api/setting/ratio_setting"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -213,6 +215,56 @@ func TestGetUserModelsFiltersByRequestedGroup(t *testing.T) {
|
||||
require.Empty(t, decodeUserModelsResponse(t, vipRecorder))
|
||||
}
|
||||
|
||||
func TestGetUserModelsExpandsAutoGroupsInConfiguredOrder(t *testing.T) {
|
||||
originalAutoGroups := setting.AutoGroups2JsonString()
|
||||
originalUsableGroups := setting.UserUsableGroups2JSONString()
|
||||
originalSpecialGroups := ratio_setting.GetGroupRatioSetting().GroupSpecialUsableGroup.ReadAll()
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, setting.UpdateAutoGroupsByJsonString(originalAutoGroups))
|
||||
require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(originalUsableGroups))
|
||||
specialGroups := ratio_setting.GetGroupRatioSetting().GroupSpecialUsableGroup
|
||||
specialGroups.Clear()
|
||||
specialGroups.AddAll(originalSpecialGroups)
|
||||
})
|
||||
|
||||
require.NoError(t, setting.UpdateAutoGroupsByJsonString(`["vip","default","unavailable"]`))
|
||||
require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(`{"auto":"自动分组","default":"默认分组","unavailable":"不可用分组"}`))
|
||||
specialGroups := ratio_setting.GetGroupRatioSetting().GroupSpecialUsableGroup
|
||||
specialGroups.Clear()
|
||||
specialGroups.Set("default", map[string]string{
|
||||
"+:vip": "VIP 分组",
|
||||
"-:unavailable": "",
|
||||
})
|
||||
|
||||
db := setupModelListControllerTestDB(t)
|
||||
require.NoError(t, db.Create(&model.User{
|
||||
Id: 1003,
|
||||
Username: "playground-auto-model-user",
|
||||
Password: "password",
|
||||
Group: "default",
|
||||
Status: common.UserStatusEnabled,
|
||||
}).Error)
|
||||
require.NoError(t, db.Create(&[]model.Ability{
|
||||
{Group: "vip", Model: "zz-vip-model", ChannelId: 1, Enabled: true},
|
||||
{Group: "vip", Model: "zz-shared-model", ChannelId: 1, Enabled: true},
|
||||
{Group: "default", Model: "zz-default-model", ChannelId: 1, Enabled: true},
|
||||
{Group: "default", Model: "zz-shared-model", ChannelId: 2, Enabled: true},
|
||||
{Group: "unavailable", Model: "zz-unavailable-model", ChannelId: 1, Enabled: true},
|
||||
}).Error)
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
context, _ := gin.CreateTestContext(recorder)
|
||||
context.Request = httptest.NewRequest(http.MethodGet, "/api/user/models?group=auto", nil)
|
||||
context.Set("id", 1003)
|
||||
|
||||
GetUserModels(context)
|
||||
|
||||
models := decodeUserModelsResponse(t, recorder)
|
||||
require.Len(t, models, 3)
|
||||
assert.ElementsMatch(t, []string{"zz-vip-model", "zz-shared-model"}, models[:2])
|
||||
assert.Equal(t, "zz-default-model", models[2])
|
||||
}
|
||||
|
||||
func TestListModelsIncludesTieredBillingModel(t *testing.T) {
|
||||
withSelfUseModeDisabled(t)
|
||||
withTieredBillingConfig(t, map[string]string{
|
||||
|
||||
+13
-25
@@ -647,38 +647,26 @@ func GetUserModels(c *gin.Context) {
|
||||
}
|
||||
groups := service.GetUserUsableGroups(user.Group)
|
||||
group := c.Query("group")
|
||||
if group != "" {
|
||||
if _, ok := groups[group]; !ok {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "",
|
||||
"data": []string{},
|
||||
})
|
||||
return
|
||||
var groupsToQuery []string
|
||||
switch {
|
||||
case group == "":
|
||||
for g := range groups {
|
||||
groupsToQuery = append(groupsToQuery, g)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "",
|
||||
"data": model.GetGroupEnabledModels(group),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
var models []string
|
||||
for group := range groups {
|
||||
for _, g := range model.GetGroupEnabledModels(group) {
|
||||
if !common.StringsContains(models, g) {
|
||||
models = append(models, g)
|
||||
}
|
||||
case group == "auto":
|
||||
if _, ok := groups[group]; ok {
|
||||
groupsToQuery = service.GetUserAutoGroup(user.Group)
|
||||
}
|
||||
default:
|
||||
if _, ok := groups[group]; ok {
|
||||
groupsToQuery = []string{group}
|
||||
}
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "",
|
||||
"data": models,
|
||||
"data": service.GetGroupsEnabledModels(groupsToQuery),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
func UpdateUser(c *gin.Context) {
|
||||
|
||||
@@ -3,6 +3,7 @@ package service
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/QuantumNous/new-api/model"
|
||||
"github.com/QuantumNous/new-api/setting"
|
||||
"github.com/QuantumNous/new-api/setting/ratio_setting"
|
||||
)
|
||||
@@ -53,6 +54,21 @@ func GetUserAutoGroup(userGroup string) []string {
|
||||
return autoGroups
|
||||
}
|
||||
|
||||
// GetGroupsEnabledModels 按 groups 顺序获取各分组启用的模型并去重
|
||||
func GetGroupsEnabledModels(groups []string) []string {
|
||||
seen := make(map[string]struct{})
|
||||
models := make([]string, 0)
|
||||
for _, group := range groups {
|
||||
for _, modelName := range model.GetGroupEnabledModels(group) {
|
||||
if _, ok := seen[modelName]; !ok {
|
||||
seen[modelName] = struct{}{}
|
||||
models = append(models, modelName)
|
||||
}
|
||||
}
|
||||
}
|
||||
return models
|
||||
}
|
||||
|
||||
// GetUserGroupRatio 获取用户使用某个分组的倍率
|
||||
// userGroup 用户分组
|
||||
// group 需要获取倍率的分组
|
||||
|
||||
Reference in New Issue
Block a user