Merge pull request #3104 from CoolCoolTomato/fix/api-key-exclusive-group-auth

fix(bug): Fix the bug related to unauthorized use of groups
This commit is contained in:
Wesley Liddick
2026-06-08 08:58:44 +08:00
committed by GitHub
6 changed files with 126 additions and 9 deletions
+17 -1
View File
@@ -107,7 +107,11 @@ func (r *apiKeyRepository) GetKeyAndOwnerID(ctx context.Context, id int64) (stri
func (r *apiKeyRepository) GetByKey(ctx context.Context, key string) (*service.APIKey, error) {
m, err := r.activeQuery().
Where(apikey.KeyEQ(key)).
WithUser().
WithUser(func(q *dbent.UserQuery) {
q.WithAllowedGroups(func(gq *dbent.GroupQuery) {
gq.Select(group.FieldID)
})
}).
WithGroup().
Only(ctx)
if err != nil {
@@ -156,12 +160,16 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se
user.FieldLastActiveAt,
user.FieldRpmLimit,
)
q.WithAllowedGroups(func(gq *dbent.GroupQuery) {
gq.Select(group.FieldID)
})
}).
WithGroup(func(q *dbent.GroupQuery) {
q.Select(
group.FieldID,
group.FieldName,
group.FieldPlatform,
group.FieldIsExclusive,
group.FieldStatus,
group.FieldSubscriptionType,
group.FieldRateMultiplier,
@@ -716,6 +724,14 @@ func apiKeyEntityToService(m *dbent.APIKey) *service.APIKey {
}
if m.Edges.User != nil {
out.User = userEntityToService(m.Edges.User)
if allowed := m.Edges.User.Edges.AllowedGroups; len(allowed) > 0 {
out.User.AllowedGroups = make([]int64, 0, len(allowed))
for _, g := range allowed {
if g != nil {
out.User.AllowedGroups = append(out.User.AllowedGroups, g.ID)
}
}
}
}
if m.Edges.Group != nil {
out.Group = groupEntityToService(m.Edges.Group)
@@ -119,6 +119,9 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
if abortIfAPIKeyGroupUnavailable(c, apiKey) {
return
}
if abortIfAPIKeyGroupNotAllowed(c, apiKey) {
return
}
// ── 4. SimpleMode → early return ─────────────────────────────
@@ -292,6 +295,26 @@ func abortIfAPIKeyGroupUnavailable(c *gin.Context, apiKey *service.APIKey) bool
return true
}
func abortIfAPIKeyGroupNotAllowed(c *gin.Context, apiKey *service.APIKey) bool {
if validateAPIKeyGroupAllowed(apiKey) {
return false
}
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonAPIKeyGroupUnavailable)
AbortWithError(c, 403, "GROUP_NOT_ALLOWED", "API Key 所属专属分组不再允许当前用户使用")
return true
}
func validateAPIKeyGroupAllowed(apiKey *service.APIKey) bool {
if apiKey == nil || apiKey.GroupID == nil || apiKey.User == nil || apiKey.Group == nil {
return true
}
group := apiKey.Group
if group.IsSubscriptionType() {
return true
}
return apiKey.User.CanBindGroup(group.ID, group.IsExclusive)
}
func validateAPIKeyGroupAvailable(apiKey *service.APIKey) (string, string, bool) {
if apiKey == nil || apiKey.GroupID == nil {
return "", "", true
@@ -235,6 +235,57 @@ func TestAPIKeyAuthSetsGroupContext(t *testing.T) {
require.Equal(t, http.StatusOK, w.Code)
}
func TestAPIKeyAuthRejectsExclusiveGroupWhenUserNoLongerAllowed(t *testing.T) {
gin.SetMode(gin.TestMode)
group := &service.Group{
ID: 202,
Name: "exclusive",
Status: service.StatusActive,
IsExclusive: true,
Hydrated: true,
}
user := &service.User{
ID: 7,
Role: service.RoleUser,
Status: service.StatusActive,
Balance: 10,
Concurrency: 3,
AllowedGroups: []int64{},
}
apiKey := &service.APIKey{
ID: 100,
UserID: user.ID,
Key: "test-key",
Status: service.StatusActive,
User: user,
Group: group,
}
apiKey.GroupID = &group.ID
apiKeyRepo := &stubApiKeyRepo{
getByKey: func(ctx context.Context, key string) (*service.APIKey, error) {
if key != apiKey.Key {
return nil, service.ErrAPIKeyNotFound
}
clone := *apiKey
return &clone, nil
},
}
cfg := &config.Config{RunMode: config.RunModeSimple}
apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
router := newAuthTestRouter(apiKeyService, nil, cfg)
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/t", nil)
req.Header.Set("x-api-key", apiKey.Key)
router.ServeHTTP(w, req)
require.Equal(t, http.StatusForbidden, w.Code)
require.Contains(t, w.Body.String(), "GROUP_NOT_ALLOWED")
}
func TestAPIKeyAuthOverwritesInvalidContextGroup(t *testing.T) {
gin.SetMode(gin.TestMode)
+23 -2
View File
@@ -749,6 +749,7 @@ func (s *adminServiceImpl) UpdateUser(ctx context.Context, id int64, input *Upda
oldStatus := user.Status
oldRole := user.Role
oldRPMLimit := user.RPMLimit
oldAllowedGroups := append([]int64(nil), user.AllowedGroups...)
if input.Email != "" {
user.Email = input.Email
@@ -795,8 +796,8 @@ func (s *adminServiceImpl) UpdateUser(ctx context.Context, id int64, input *Upda
if s.authCacheInvalidator != nil {
// RPMLimit 直接参与 billing_cache_service.checkRPM 的三级级联,
// 不失效缓存会让修改在一个 L2 TTL 内失去效果。
if user.Concurrency != oldConcurrency || user.Status != oldStatus || user.Role != oldRole || user.RPMLimit != oldRPMLimit {
// allowed_groups 参与 API Key 专属分组授权判断;不失效缓存会让修改在一个 L2 TTL 内失去效果。
if user.Concurrency != oldConcurrency || user.Status != oldStatus || user.Role != oldRole || user.RPMLimit != oldRPMLimit || !sameInt64Set(user.AllowedGroups, oldAllowedGroups) {
s.authCacheInvalidator.InvalidateAuthCacheByUserID(ctx, user.ID)
}
}
@@ -825,6 +826,26 @@ func (s *adminServiceImpl) UpdateUser(ctx context.Context, id int64, input *Upda
return user, nil
}
func sameInt64Set(a, b []int64) bool {
if len(a) != len(b) {
return false
}
if len(a) == 0 {
return true
}
counts := make(map[int64]int, len(a))
for _, v := range a {
counts[v]++
}
for _, v := range b {
if counts[v] == 0 {
return false
}
counts[v]--
}
return true
}
func (s *adminServiceImpl) DeleteUser(ctx context.Context, id int64) error {
// Protect admin users: cannot delete admin accounts
user, err := s.userRepo.GetByID(ctx, id)
@@ -30,11 +30,12 @@ type APIKeyAuthSnapshot struct {
// APIKeyAuthUserSnapshot 用户快照
type APIKeyAuthUserSnapshot struct {
ID int64 `json:"id"`
Status string `json:"status"`
Role string `json:"role"`
Balance float64 `json:"balance"`
Concurrency int `json:"concurrency"`
ID int64 `json:"id"`
Status string `json:"status"`
Role string `json:"role"`
Balance float64 `json:"balance"`
Concurrency int `json:"concurrency"`
AllowedGroups []int64 `json:"allowed_groups,omitempty"`
// Balance notification fields (required for CheckBalanceAfterDeduction)
Email string `json:"email"`
@@ -58,6 +59,7 @@ type APIKeyAuthGroupSnapshot struct {
ID int64 `json:"id"`
Name string `json:"name"`
Platform string `json:"platform"`
IsExclusive bool `json:"is_exclusive"`
Status string `json:"status"`
SubscriptionType string `json:"subscription_type"`
RateMultiplier float64 `json:"rate_multiplier"`
@@ -14,7 +14,7 @@ import (
"github.com/dgraph-io/ristretto"
)
const apiKeyAuthSnapshotVersion = 11 // v11: reload snapshots for custom models_list_config
const apiKeyAuthSnapshotVersion = 12 // v12: include exclusive group authorization fields
type apiKeyAuthCacheConfig struct {
l1Size int
@@ -226,6 +226,7 @@ func (s *APIKeyService) snapshotFromAPIKey(ctx context.Context, apiKey *APIKey)
Role: apiKey.User.Role,
Balance: apiKey.User.Balance,
Concurrency: apiKey.User.Concurrency,
AllowedGroups: apiKey.User.AllowedGroups,
Email: apiKey.User.Email,
Username: apiKey.User.Username,
BalanceNotifyEnabled: apiKey.User.BalanceNotifyEnabled,
@@ -250,6 +251,7 @@ func (s *APIKeyService) snapshotFromAPIKey(ctx context.Context, apiKey *APIKey)
ID: apiKey.Group.ID,
Name: apiKey.Group.Name,
Platform: apiKey.Group.Platform,
IsExclusive: apiKey.Group.IsExclusive,
Status: apiKey.Group.Status,
SubscriptionType: apiKey.Group.SubscriptionType,
RateMultiplier: apiKey.Group.RateMultiplier,
@@ -304,6 +306,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho
Role: snapshot.User.Role,
Balance: snapshot.User.Balance,
Concurrency: snapshot.User.Concurrency,
AllowedGroups: snapshot.User.AllowedGroups,
Email: snapshot.User.Email,
Username: snapshot.User.Username,
BalanceNotifyEnabled: snapshot.User.BalanceNotifyEnabled,
@@ -320,6 +323,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho
ID: snapshot.Group.ID,
Name: snapshot.Group.Name,
Platform: snapshot.Group.Platform,
IsExclusive: snapshot.Group.IsExclusive,
Status: snapshot.Group.Status,
Hydrated: true,
SubscriptionType: snapshot.Group.SubscriptionType,