Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 56 additions & 0 deletions common/group.go
Original file line number Diff line number Diff line change
@@ -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, ",")
}
42 changes: 42 additions & 0 deletions common/group_test.go
Original file line number Diff line number Diff line change
@@ -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(""))
}
14 changes: 9 additions & 5 deletions common/topup-ratio.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
2 changes: 1 addition & 1 deletion controller/channel.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
10 changes: 7 additions & 3 deletions controller/model.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand Down
11 changes: 11 additions & 0 deletions controller/user.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 :)
}
Expand Down
6 changes: 4 additions & 2 deletions middleware/auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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 {
Expand Down
4 changes: 3 additions & 1 deletion middleware/distributor.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
6 changes: 3 additions & 3 deletions middleware/model-rate-limit.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks, fixed in d52111b.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in d52111b: the ModelRequestRateLimit lookup now applies the same resolution — token group first, then common.PrimaryGroup(ContextKeyUserGroup) — so a multi-group user matches the configured per-group rate limit instead of silently falling back to global limits.

}
modelName := getRateLimitModelName(c)

Expand Down Expand Up @@ -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))
}

//获取分组的限流配置
Expand Down
8 changes: 5 additions & 3 deletions model/channel.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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").
Expand Down
2 changes: 1 addition & 1 deletion model/subscription.go
Original file line number Diff line number Diff line change
Expand Up @@ -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:''"`
Expand Down
6 changes: 4 additions & 2 deletions model/user.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"` // 邀请剩余额度
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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{}{
Expand Down
69 changes: 69 additions & 0 deletions model/user_group_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
9 changes: 7 additions & 2 deletions relay/common/relay_info.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading
Loading