diff --git a/common/group.go b/common/group.go new file mode 100644 index 000000000000..1ccfa8d75495 --- /dev/null +++ b/common/group.go @@ -0,0 +1,56 @@ +package common + +import "strings" + +const ( + // MaxGroupNameLength is the maximum length of a single group name in a + // comma-separated group list. + MaxGroupNameLength = 64 + // MaxUserGroupColumnLength matches the users.group column width + // (varchar(1024)) that stores the comma-separated multi-group list. + MaxUserGroupColumnLength = 1024 +) + +// SplitGroupList splits a comma-separated group list (e.g. a multi-group user +// or channel) into trimmed, non-empty, de-duplicated group names, preserving +// the original order. +func SplitGroupList(s string) []string { + if strings.TrimSpace(s) == "" { + return nil + } + parts := strings.Split(s, ",") + groups := make([]string, 0, len(parts)) + seen := make(map[string]struct{}, len(parts)) + for _, part := range parts { + group := strings.TrimSpace(part) + if group == "" { + continue + } + if _, ok := seen[group]; ok { + continue + } + seen[group] = struct{}{} + groups = append(groups, group) + } + return groups +} + +// PrimaryGroup returns the first group of a comma-separated list, falling back +// to "default" when the list contains no valid group. +func PrimaryGroup(s string) string { + if groups := SplitGroupList(s); len(groups) > 0 { + return groups[0] + } + return "default" +} + +// NormalizeGroupList canonicalizes a comma-separated group list by trimming +// whitespace and removing empty duplicates. It returns "default" when nothing +// valid remains so a user always keeps at least one group. +func NormalizeGroupList(s string) string { + groups := SplitGroupList(s) + if len(groups) == 0 { + return "default" + } + return strings.Join(groups, ",") +} diff --git a/common/group_test.go b/common/group_test.go new file mode 100644 index 000000000000..7dc1d8333446 --- /dev/null +++ b/common/group_test.go @@ -0,0 +1,42 @@ +package common + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestSplitGroupList(t *testing.T) { + tests := []struct { + name string + input string + want []string + }{ + {"empty", "", nil}, + {"spaces only", " , ", []string{}}, + {"single", "default", []string{"default"}}, + {"multi", "group0,group1", []string{"group0", "group1"}}, + {"trims whitespace", " group0 , group1 ", []string{"group0", "group1"}}, + {"drops empties", "group0,,group1,", []string{"group0", "group1"}}, + {"dedupes preserving order", "b,a,b", []string{"b", "a"}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, SplitGroupList(tt.input)) + }) + } +} + +func TestPrimaryGroup(t *testing.T) { + assert.Equal(t, "group0", PrimaryGroup("group0,group1")) + assert.Equal(t, "group0", PrimaryGroup(" group0 , group1 ")) + assert.Equal(t, "group0", PrimaryGroup("group0")) + assert.Equal(t, "default", PrimaryGroup("")) + assert.Equal(t, "default", PrimaryGroup(" , ")) +} + +func TestNormalizeGroupList(t *testing.T) { + assert.Equal(t, "group0,group1", NormalizeGroupList(" group0 , , group1 ")) + assert.Equal(t, "group0", NormalizeGroupList("group0,group0")) + assert.Equal(t, "default", NormalizeGroupList("")) +} diff --git a/common/topup-ratio.go b/common/topup-ratio.go index 2b60cde7d169..e432cbe6ad50 100644 --- a/common/topup-ratio.go +++ b/common/topup-ratio.go @@ -29,13 +29,17 @@ func UpdateTopupGroupRatioByJSONString(jsonStr string) error { return json.Unmarshal([]byte(jsonStr), &topupGroupRatio) } +// GetTopupGroupRatio returns the topup ratio for the given user group. name +// may be a comma-separated multi-group list; groups are checked in list order +// and the first configured entry wins. func GetTopupGroupRatio(name string) float64 { topupGroupRatioMutex.RLock() defer topupGroupRatioMutex.RUnlock() - ratio, ok := topupGroupRatio[name] - if !ok { - SysError("topup group ratio not found: " + name) - return 1 + for _, group := range SplitGroupList(name) { + if ratio, ok := topupGroupRatio[group]; ok { + return ratio + } } - return ratio + SysError("topup group ratio not found: " + name) + return 1 } diff --git a/controller/channel.go b/controller/channel.go index 69c1561f2f3d..1b5aa360366a 100644 --- a/controller/channel.go +++ b/controller/channel.go @@ -83,7 +83,7 @@ func applyChannelStatusFilter(query *gorm.DB, statusFilter int) *gorm.DB { func buildChannelListQuery(group string, statusFilter int, typeFilter int) *gorm.DB { query := model.DB.Model(&model.Channel{}) - query = model.ApplyChannelGroupFilter(query, group) + query = model.ApplyGroupContainsFilter(query, group) query = applyChannelStatusFilter(query, statusFilter) if typeFilter >= 0 { query = query.Where("type = ?", typeFilter) diff --git a/controller/model.go b/controller/model.go index 1d759301bc7e..8b821436243a 100644 --- a/controller/model.go +++ b/controller/model.go @@ -195,14 +195,18 @@ func getModelListGroups(c *gin.Context) (modelListGroups, error) { }, nil } - group := userGroup if tokenGroup != "" { - group = tokenGroup + return modelListGroups{ + userGroup: userGroup, + tokenGroup: tokenGroup, + ownerGroups: []string{tokenGroup}, + }, nil } + // 多分组用户返回其所有分组下启用的模型 return modelListGroups{ userGroup: userGroup, tokenGroup: tokenGroup, - ownerGroups: []string{group}, + ownerGroups: common.SplitGroupList(userGroup), }, nil } diff --git a/controller/user.go b/controller/user.go index 9b8d931ec1f8..fe7ed13bdfb0 100644 --- a/controller/user.go +++ b/controller/user.go @@ -674,6 +674,17 @@ func UpdateUser(c *gin.Context) { common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } + updatedUser.Group = common.NormalizeGroupList(updatedUser.Group) + for _, group := range common.SplitGroupList(updatedUser.Group) { + if len(group) > common.MaxGroupNameLength { + common.ApiErrorI18n(c, i18n.MsgInvalidParams) + return + } + } + if len(updatedUser.Group) > common.MaxUserGroupColumnLength { + common.ApiErrorI18n(c, i18n.MsgInvalidParams) + return + } if updatedUser.Password == "" { updatedUser.Password = "$I_LOVE_U" // make Validator happy :) } diff --git a/middleware/auth.go b/middleware/auth.go index d20385309886..91847e7bcbcb 100644 --- a/middleware/auth.go +++ b/middleware/auth.go @@ -460,6 +460,8 @@ func TokenAuth() func(c *gin.Context) { userCache.WriteContext(c) userGroup := userCache.Group + // 多分组用户(逗号分隔)默认使用第一个分组,令牌可指定其他可用分组 + usingGroup := common.PrimaryGroup(userGroup) tokenGroup := token.Group if tokenGroup != "" { // check common.UserUsableGroups[userGroup] @@ -474,9 +476,9 @@ func TokenAuth() func(c *gin.Context) { return } } - userGroup = tokenGroup + usingGroup = tokenGroup } - common.SetContextKey(c, constant.ContextKeyUsingGroup, userGroup) + common.SetContextKey(c, constant.ContextKeyUsingGroup, usingGroup) err = SetupContextForToken(c, token, parts...) if err != nil { diff --git a/middleware/distributor.go b/middleware/distributor.go index a17f9c50f569..b52963fdf95d 100644 --- a/middleware/distributor.go +++ b/middleware/distributor.go @@ -93,7 +93,9 @@ func Distribute() func(c *gin.Context) { return } if playgroundRequest.Group != "" { - if !service.GroupInUserUsableGroups(usingGroup, playgroundRequest.Group) && playgroundRequest.Group != usingGroup { + // 多分组用户按其完整分组集合校验可用性 + userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup) + if !service.GroupInUserUsableGroups(userGroup, playgroundRequest.Group) && playgroundRequest.Group != usingGroup { abortWithOpenAiMessage(c, http.StatusForbidden, i18n.T(c, i18n.MsgDistributorGroupAccessDenied)) return } diff --git a/middleware/model-rate-limit.go b/middleware/model-rate-limit.go index da4a1dba8ec6..c2e117cc46e2 100644 --- a/middleware/model-rate-limit.go +++ b/middleware/model-rate-limit.go @@ -49,7 +49,7 @@ func recordRateLimitErrorLog(c *gin.Context, statusCode int, message string) { tokenId := c.GetInt("token_id") group := common.GetContextKeyString(c, constant.ContextKeyTokenGroup) if group == "" { - group = common.GetContextKeyString(c, constant.ContextKeyUserGroup) + group = common.PrimaryGroup(common.GetContextKeyString(c, constant.ContextKeyUserGroup)) } modelName := getRateLimitModelName(c) @@ -226,10 +226,10 @@ func ModelRequestRateLimit() func(c *gin.Context) { totalMaxCount := setting.ModelRequestRateLimitCount successMaxCount := setting.ModelRequestRateLimitSuccessCount - // 获取分组 + // 获取分组;多分组用户取主分组,避免逗号列表无法命中分组限流配置 group := common.GetContextKeyString(c, constant.ContextKeyTokenGroup) if group == "" { - group = common.GetContextKeyString(c, constant.ContextKeyUserGroup) + group = common.PrimaryGroup(common.GetContextKeyString(c, constant.ContextKeyUserGroup)) } //获取分组的限流配置 diff --git a/model/channel.go b/model/channel.go index 2dd03fc4b895..bd4b95ec8361 100644 --- a/model/channel.go +++ b/model/channel.go @@ -166,7 +166,9 @@ func channelGroupFilterPattern(group string) string { return "%," + group + ",%" } -func ApplyChannelGroupFilter(query *gorm.DB, group string) *gorm.DB { +// ApplyGroupContainsFilter filters rows whose comma-separated group column +// (channels and multi-group users) contains the given group. +func ApplyGroupContainsFilter(query *gorm.DB, group string) *gorm.DB { group = NormalizeChannelGroupFilter(group) if group == "" { return query @@ -444,7 +446,7 @@ func SearchChannels(keyword string, group string, model string, idSort bool, sor // 构造WHERE子句 whereClause := "(id = ? OR name LIKE ? OR " + commonKeyCol + " = ? OR " + baseURLCol + " LIKE ?) AND " + modelsCol + " LIKE ?" args := []any{common.String2Int(keyword), "%" + keyword + "%", keyword, "%" + keyword + "%", "%" + model + "%"} - baseQuery = ApplyChannelGroupFilter(baseQuery.Where(whereClause, args...), group) + baseQuery = ApplyGroupContainsFilter(baseQuery.Where(whereClause, args...), group) // 执行查询 err := order.Apply(baseQuery).Find(&channels).Error @@ -1042,7 +1044,7 @@ func SearchTags(keyword string, group string, model string, idSort bool) ([]*str // 构造WHERE子句 whereClause := "(id = ? OR name LIKE ? OR " + commonKeyCol + " = ? OR " + baseURLCol + " LIKE ?) AND " + modelsCol + " LIKE ?" args := []any{common.String2Int(keyword), "%" + keyword + "%", keyword, "%" + keyword + "%", "%" + model + "%"} - baseQuery = ApplyChannelGroupFilter(baseQuery.Where(whereClause, args...), group) + baseQuery = ApplyGroupContainsFilter(baseQuery.Where(whereClause, args...), group) subQuery := baseQuery. Select("tag"). diff --git a/model/subscription.go b/model/subscription.go index c89e63bf0532..ff4f3f549d58 100644 --- a/model/subscription.go +++ b/model/subscription.go @@ -268,7 +268,7 @@ type UserSubscription struct { NextResetTime int64 `json:"next_reset_time" gorm:"type:bigint;default:0;index"` UpgradeGroup string `json:"upgrade_group" gorm:"type:varchar(64);default:''"` - PrevUserGroup string `json:"prev_user_group" gorm:"type:varchar(64);default:''"` + PrevUserGroup string `json:"prev_user_group" gorm:"type:varchar(1024);default:''"` // 可能是多分组用户的逗号列表快照 // Downgrade target group on expiry (snapshot from plan; empty = revert to PrevUserGroup) DowngradeGroup string `json:"downgrade_group" gorm:"type:varchar(64);default:''"` diff --git a/model/user.go b/model/user.go index 1d38ea392ad5..16a77adf0575 100644 --- a/model/user.go +++ b/model/user.go @@ -95,7 +95,7 @@ type User struct { Quota int `json:"quota" gorm:"type:int;default:0"` UsedQuota int `json:"used_quota" gorm:"type:int;default:0;column:used_quota"` // used quota RequestCount int `json:"request_count" gorm:"type:int;default:0;"` // request number - Group string `json:"group" gorm:"type:varchar(64);default:'default'"` + Group string `json:"group" gorm:"type:varchar(1024);default:'default'"` // 逗号分隔的多分组,第一个为主分组 AffCode string `json:"aff_code" gorm:"type:varchar(32);column:aff_code;uniqueIndex"` AffCount int `json:"aff_count" gorm:"type:int;default:0;column:aff_count"` AffQuota int `json:"aff_quota" gorm:"type:int;default:0;column:aff_quota"` // 邀请剩余额度 @@ -454,7 +454,8 @@ func SearchUsers(keyword string, group string, role *int, status *int, startIdx query = query.Where("("+likeCondition+")", likeArgs...) if group != "" { - query = query.Where(commonGroupCol+" = ?", group) + // 多分组用户按“包含该分组”匹配 + query = ApplyGroupContainsFilter(query, group) } if role != nil { query = query.Where("role = ?", *role) @@ -835,6 +836,7 @@ func (user *User) EditWithTx(tx *gorm.DB, updatePassword bool) error { return err } } + user.Group = common.NormalizeGroupList(user.Group) newUser := *user updates := map[string]interface{}{ diff --git a/model/user_group_test.go b/model/user_group_test.go new file mode 100644 index 000000000000..3be090601be3 --- /dev/null +++ b/model/user_group_test.go @@ -0,0 +1,69 @@ +package model + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func insertUsersForGroupFilterTest(t *testing.T) { + t.Helper() + truncateTables(t) + users := []*User{ + {Username: "single", Password: "password123", Group: "group0", AffCode: "aff-s0"}, + {Username: "multi", Password: "password123", Group: "group0,group1", AffCode: "aff-m0"}, + {Username: "other", Password: "password123", Group: "group10", AffCode: "aff-o0"}, + } + for _, user := range users { + require.NoError(t, DB.Create(user).Error) + } +} + +func TestSearchUsersGroupFilterMatchesMultiGroupUsers(t *testing.T) { + insertUsersForGroupFilterTest(t) + + users, total, err := SearchUsers("", "group0", nil, nil, 0, 10) + require.NoError(t, err) + assert.Equal(t, int64(2), total) + assert.ElementsMatch(t, []string{"single", "multi"}, usernames(users)) + + // 按列表中任意一个分组过滤都能命中多分组用户; + // 同时验证不会用子串误匹配:group1 不命中 group10 + users, total, err = SearchUsers("", "group1", nil, nil, 0, 10) + require.NoError(t, err) + assert.Equal(t, int64(1), total) + require.Len(t, users, 1) + assert.Equal(t, "multi", users[0].Username) + + _, total, err = SearchUsers("", "group10", nil, nil, 0, 10) + require.NoError(t, err) + assert.Equal(t, int64(1), total) +} + +func usernames(users []*User) []string { + names := make([]string, 0, len(users)) + for _, user := range users { + names = append(names, user.Username) + } + return names +} + +func TestEditWithTxNormalizesGroupList(t *testing.T) { + insertUsersForGroupFilterTest(t) + var multi User + require.NoError(t, DB.Where("username = ?", "multi").First(&multi).Error) + + multi.Group = " group1 , , group0 , group1 " + require.NoError(t, multi.EditWithTx(DB, false)) + assert.Equal(t, "group1,group0", multi.Group) + + var reloaded User + require.NoError(t, DB.Where("username = ?", "multi").First(&reloaded).Error) + assert.Equal(t, "group1,group0", reloaded.Group) + + // 空分组规范化为 default + multi.Group = " , " + require.NoError(t, multi.EditWithTx(DB, false)) + assert.Equal(t, "default", multi.Group) +} diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index b0bb19bdca3b..a19901890ae6 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -477,9 +477,14 @@ func genBaseRelayInfo(c *gin.Context, request dto.Request) *RelayInfo { //paramOverride := common.GetContextKeyStringMap(c, constant.ContextKeyChannelParamOverride) tokenGroup := common.GetContextKeyString(c, constant.ContextKeyTokenGroup) - // 当令牌分组为空时,表示使用用户分组 + // 当令牌分组为空时,表示使用用户分组;多分组用户取实际使用的单一分组 + // (优先取 distributor 解析出的 UsingGroup,否则取用户的主分组), + // 该值会进入渠道选择的精确匹配,不能是逗号列表。 if tokenGroup == "" { - tokenGroup = common.GetContextKeyString(c, constant.ContextKeyUserGroup) + tokenGroup = common.GetContextKeyString(c, constant.ContextKeyUsingGroup) + } + if tokenGroup == "" { + tokenGroup = common.PrimaryGroup(common.GetContextKeyString(c, constant.ContextKeyUserGroup)) } startTime := common.GetContextKeyTime(c, constant.ContextKeyRequestStartTime) diff --git a/service/group.go b/service/group.go index e792083e4a09..25f62d52fafc 100644 --- a/service/group.go +++ b/service/group.go @@ -11,30 +11,39 @@ import ( "github.com/gin-gonic/gin" ) +// GetUserUsableGroups returns the usable-group set for a user. userGroup may +// be a comma-separated multi-group list; special usable rules are applied per +// group in list order (later groups can remove groups added by earlier ones), +// then every group the user belongs to is re-added so owned groups can never +// be removed by special rules. func GetUserUsableGroups(userGroup string) map[string]string { groupsCopy := setting.GetUserUsableGroupsCopy() - if userGroup != "" { - specialSettings, b := ratio_setting.GetGroupRatioSetting().GroupSpecialUsableGroup.Get(userGroup) - if b { - // 处理特殊可用分组 - for specialGroup, desc := range specialSettings { - if strings.HasPrefix(specialGroup, "-:") { - // 移除分组 - groupToRemove := strings.TrimPrefix(specialGroup, "-:") - delete(groupsCopy, groupToRemove) - } else if strings.HasPrefix(specialGroup, "+:") { - // 添加分组 - groupToAdd := strings.TrimPrefix(specialGroup, "+:") - groupsCopy[groupToAdd] = desc - } else { - // 直接添加分组 - groupsCopy[specialGroup] = desc - } + groups := common.SplitGroupList(userGroup) + for _, group := range groups { + specialSettings, b := ratio_setting.GetGroupRatioSetting().GroupSpecialUsableGroup.Get(group) + if !b { + continue + } + // 处理特殊可用分组 + for specialGroup, desc := range specialSettings { + if strings.HasPrefix(specialGroup, "-:") { + // 移除分组 + groupToRemove := strings.TrimPrefix(specialGroup, "-:") + delete(groupsCopy, groupToRemove) + } else if strings.HasPrefix(specialGroup, "+:") { + // 添加分组 + groupToAdd := strings.TrimPrefix(specialGroup, "+:") + groupsCopy[groupToAdd] = desc + } else { + // 直接添加分组 + groupsCopy[specialGroup] = desc } } - // 如果userGroup不在UserUsableGroups中,返回UserUsableGroups + userGroup - if _, ok := groupsCopy[userGroup]; !ok { - groupsCopy[userGroup] = "用户分组" + } + // 所属分组本身始终可选,不能被特殊规则移除 + for _, group := range groups { + if _, ok := groupsCopy[group]; !ok { + groupsCopy[group] = "用户分组" } } return groupsCopy diff --git a/service/group_multi_test.go b/service/group_multi_test.go new file mode 100644 index 000000000000..3443a7a4e5ed --- /dev/null +++ b/service/group_multi_test.go @@ -0,0 +1,121 @@ +package service + +import ( + "testing" + + "github.com/QuantumNous/new-api/setting" + "github.com/QuantumNous/new-api/setting/ratio_setting" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func configureMultiGroupTest(t *testing.T) { + t.Helper() + originalUsableGroups := setting.UserUsableGroups2JSONString() + originalRatios := ratio_setting.GroupRatio2JSONString() + originalGroupGroupRatios := ratio_setting.GroupGroupRatio2JSONString() + specialUsable := ratio_setting.GetGroupRatioSetting().GroupSpecialUsableGroup + originalSpecialUsable := specialUsable.ReadAll() + + require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(`{"default":"默认分组","vip":"VIP分组"}`)) + require.NoError(t, ratio_setting.UpdateGroupRatioByJSONString(`{"default":1,"vip":1,"group0":1,"group1":1,"group2":1,"premium":1}`)) + + t.Cleanup(func() { + require.NoError(t, setting.UpdateUserUsableGroupsByJSONString(originalUsableGroups)) + require.NoError(t, ratio_setting.UpdateGroupRatioByJSONString(originalRatios)) + require.NoError(t, ratio_setting.UpdateGroupGroupRatioByJSONString(originalGroupGroupRatios)) + specialUsable.Clear() + specialUsable.AddAll(originalSpecialUsable) + }) +} + +func TestGetUserUsableGroupsMultiGroupUnion(t *testing.T) { + configureMultiGroupTest(t) + + groups := GetUserUsableGroups("group0,group1") + + // 全局可用分组 + 用户所属的每个分组本身 + for _, group := range []string{"default", "vip", "group0", "group1"} { + assert.Contains(t, groups, group) + } + assert.NotContains(t, groups, "group2") +} + +func TestGetUserUsableGroupsMultiGroupSpecialRulesAppliedInOrder(t *testing.T) { + configureMultiGroupTest(t) + specialUsable := ratio_setting.GetGroupRatioSetting().GroupSpecialUsableGroup + specialUsable.Set("group0", map[string]string{"+:premium": "Premium"}) + specialUsable.Set("group1", map[string]string{"-:vip": "", "-:group0": ""}) + + groups := GetUserUsableGroups("group0,group1") + + // group0 的规则先添加 premium,group1 的规则随后移除 vip + assert.Contains(t, groups, "premium") + assert.NotContains(t, groups, "vip") + // 所属分组本身始终可选,不能被其他所属分组的特殊规则移除 + assert.Contains(t, groups, "group0") + assert.Contains(t, groups, "group1") +} + +func TestGetUserGroupRatioMultiGroupFirstMatchWins(t *testing.T) { + configureMultiGroupTest(t) + require.NoError(t, ratio_setting.UpdateGroupGroupRatioByJSONString( + `{"group1":{"vip":0.8},"group0":{"vip":0.5,"group1":1.2}}`)) + + // 两个所属分组都配置了 vip 的专属倍率时,按列表顺序取第一个(主分组优先) + ratio := GetUserGroupRatio("group0,group1", "vip") + assert.Equal(t, 0.5, ratio) + + // 仅第二个分组配置时使用第二个分组 + ratio = GetUserGroupRatio("group2,group1", "vip") + assert.Equal(t, 0.8, ratio) + + // 没有任何专属倍率时回退到全局分组倍率 + ratio = GetUserGroupRatio("group0,group1", "premium") + assert.Equal(t, 1.0, ratio) +} + +func TestGetUserGroupRatioSingleGroupUnchanged(t *testing.T) { + configureMultiGroupTest(t) + require.NoError(t, ratio_setting.UpdateGroupGroupRatioByJSONString( + `{"vip":{"default":0.9}}`)) + + assert.Equal(t, 0.9, GetUserGroupRatio("vip", "default")) + assert.Equal(t, 1.0, GetUserGroupRatio("default", "default")) +} + +func TestIsUserSelectableGroupMultiGroup(t *testing.T) { + configureMultiGroupTest(t) + require.NoError(t, setting.UpdateAutoGroupsByJsonString(`["vip","group1"]`)) + t.Cleanup(func() { + require.NoError(t, setting.UpdateAutoGroupsByJsonString("[]")) + }) + + // 所属任一分组的可选项均可用 + assert.True(t, IsUserSelectableGroup("group0,group1", "vip")) + assert.True(t, IsUserSelectableGroup("group0,group1", "group1")) + assert.False(t, IsUserSelectableGroup("group0,group1", "group2")) +} + +func TestGetUserAutoGroupMultiGroup(t *testing.T) { + configureMultiGroupTest(t) + originalAutoGroups := setting.AutoGroups2JsonString() + require.NoError(t, setting.UpdateAutoGroupsByJsonString(`["vip","group0","group1"]`)) + t.Cleanup(func() { + require.NoError(t, setting.UpdateAutoGroupsByJsonString(originalAutoGroups)) + }) + + groups := GetUserAutoGroup("group1,group0") + + // Auto 列表顺序保持不变:vip 来自全局可用分组,group0/group1 来自用户所属分组 + assert.Equal(t, []string{"vip", "group0", "group1"}, groups) +} + +func TestGetUserUsableGroupsEmptyIsGlobalOnly(t *testing.T) { + configureMultiGroupTest(t) + + groups := GetUserUsableGroups("") + delete(groups, "default") + delete(groups, "vip") + assert.Empty(t, groups) +} diff --git a/service/task_billing.go b/service/task_billing.go index 64dbdcef14a5..fc11dcc8375f 100644 --- a/service/task_billing.go +++ b/service/task_billing.go @@ -297,20 +297,26 @@ func RecalculateTaskQuotaByTokens(ctx context.Context, task *model.Task, totalTo return } - // 获取用户和组的倍率信息 + // group 为任务使用的分组,userGroup 为用户分组列表(可能是逗号分隔的多分组), + // 用户-分组专属倍率需按用户的完整分组列表按序匹配 group := task.Group - if group == "" { - user, err := model.GetUserById(task.UserId, false) - if err == nil { - group = user.Group + userGroup := "" + if user, err := model.GetUserById(task.UserId, false); err == nil { + userGroup = user.Group + if group == "" { + group = common.PrimaryGroup(userGroup) } } if group == "" { return } + if userGroup == "" { + // 用户分组查询失败时退回使用分组,保持旧的精确匹配行为 + userGroup = group + } groupRatio := ratio_setting.GetGroupRatio(group) - userGroupRatio, hasUserGroupRatio := ratio_setting.GetGroupGroupRatio(group, group) + userGroupRatio, hasUserGroupRatio := ratio_setting.GetGroupGroupRatio(userGroup, group) var finalGroupRatio float64 if hasUserGroupRatio { diff --git a/setting/ratio_setting/group_ratio.go b/setting/ratio_setting/group_ratio.go index 7d16d9283932..221de06660e8 100644 --- a/setting/ratio_setting/group_ratio.go +++ b/setting/ratio_setting/group_ratio.go @@ -85,16 +85,22 @@ func GetGroupRatio(name string) float64 { return ratio } +// GetGroupGroupRatio returns the special ratio configured for userGroup when +// requests run in usingGroup. userGroup may be a comma-separated multi-group +// list; groups are checked in list order and the first configured entry wins. func GetGroupGroupRatio(userGroup, usingGroup string) (float64, bool) { - gp, ok := groupGroupRatioMap.Get(userGroup) - if !ok { - return -1, false - } - ratio, ok := gp[usingGroup] - if !ok { - return -1, false + for _, group := range common.SplitGroupList(userGroup) { + gp, ok := groupGroupRatioMap.Get(group) + if !ok { + continue + } + ratio, ok := gp[usingGroup] + if !ok { + continue + } + return ratio, true } - return ratio, true + return -1, false } func GroupGroupRatio2JSONString() string { diff --git a/web/src/components/multi-select.tsx b/web/src/components/multi-select.tsx index ffbce5a2c713..cecfd2adcf3a 100644 --- a/web/src/components/multi-select.tsx +++ b/web/src/components/multi-select.tsx @@ -62,6 +62,10 @@ interface MultiSelectProps { id?: string /** Disable the entire control. */ disabled?: boolean + /** Forwarded to the underlying input so FormControl aria wiring reaches screen readers. */ + 'aria-describedby'?: string + /** Forwarded to the underlying input so FormControl aria wiring reaches screen readers. */ + 'aria-invalid'?: boolean | 'false' | 'true' /** * Limits rendered chips while keeping all values selected. * Hidden values remain searchable/removable from the dropdown. @@ -340,6 +344,8 @@ export function MultiSelect(props: MultiSelectProps) { [] { header: t('Group'), cell: ({ row }) => { const group = row.getValue('group') as string + const groups = parseGroupList(group) + if (groups.length === 0) { + return ( + + + + ) + } return ( - +
+ {groups.map((name) => ( + + ))} +
) }, filterFn: (row, id, value) => { - const group = String(row.getValue(id) || t('User Group')).toLowerCase() + const groups = parseGroupList(row.getValue(id) as string) const searchValue = String(value).toLowerCase() - return group.includes(searchValue) + return groups.some((group) => group.toLowerCase().includes(searchValue)) }, size: 140, meta: { mobileOrder: 30 }, diff --git a/web/src/features/users/components/users-mutate-drawer.tsx b/web/src/features/users/components/users-mutate-drawer.tsx index 8409f5721d64..8924778604cb 100644 --- a/web/src/features/users/components/users-mutate-drawer.tsx +++ b/web/src/features/users/components/users-mutate-drawer.tsx @@ -31,6 +31,7 @@ import { sideDrawerFormClassName, sideDrawerHeaderClassName, } from '@/components/drawer-layout' +import { MultiSelect } from '@/components/multi-select' import { Button } from '@/components/ui/button' import { Checkbox } from '@/components/ui/checkbox' import { @@ -71,6 +72,7 @@ import { } from '@/lib/admin-permissions' import { getCurrencyDisplay, getCurrencyLabel } from '@/lib/currency' import { formatQuota, parseQuotaFromDollars } from '@/lib/format' +import { parseGroupList } from '@/lib/group-list' import { ROLE } from '@/lib/roles' import { useAuthStore } from '@/stores/auth-store' @@ -356,37 +358,33 @@ export function UsersMutateDrawer({ ( - - {t('Group')} - - - - )} + + {t( + 'Users can belong to multiple groups; the first one is the primary group.' + )} + + + + ) + }} /> . + +For commercial licensing, please contact support@quantumnous.com +*/ +import { describe, expect, test } from 'vitest' + +import { parseGroupList } from './group-list' + +describe('parseGroupList', () => { + test('returns an empty array for missing values', () => { + expect(parseGroupList(undefined)).toEqual([]) + expect(parseGroupList(null)).toEqual([]) + expect(parseGroupList('')).toEqual([]) + }) + + test('splits a comma-separated list and trims whitespace', () => { + expect(parseGroupList('group0,group1')).toEqual(['group0', 'group1']) + expect(parseGroupList(' group0 , group1 ')).toEqual(['group0', 'group1']) + }) + + test('drops empty entries', () => { + expect(parseGroupList('group0,,group1,')).toEqual(['group0', 'group1']) + expect(parseGroupList(' , ')).toEqual([]) + }) + + test('keeps order so the first entry is the primary group', () => { + expect(parseGroupList('group1,group0')[0]).toBe('group1') + }) + + test('does not split substrings into separate groups', () => { + expect(parseGroupList('group10')).toEqual(['group10']) + }) +}) diff --git a/web/src/lib/group-list.ts b/web/src/lib/group-list.ts new file mode 100644 index 000000000000..01d092b81cba --- /dev/null +++ b/web/src/lib/group-list.ts @@ -0,0 +1,32 @@ +/* +Copyright (C) 2023-2026 QuantumNous + +This program is free software: you can redistribute it and/or modify +it under the terms of the GNU Affero General Public License as +published by the Free Software Foundation, either version 3 of the +License, or (at your option) any later version. + +This program is distributed in the hope that it will be useful, +but WITHOUT ANY WARRANTY; without even the implied warranty of +MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +GNU Affero General Public License for more details. + +You should have received a copy of the GNU Affero General Public License +along with this program. If not, see . + +For commercial licensing, please contact support@quantumnous.com +*/ + +/** + * Parse a comma-separated group list (e.g. a multi-group user) into trimmed, + * non-empty group names. The first entry is the primary group. + */ +export function parseGroupList(value?: string | null): string[] { + if (!value) { + return [] + } + return value + .split(',') + .map((group) => group.trim()) + .filter((group) => group.length > 0) +}