diff --git a/.gitignore b/.gitignore index da11257671..f87dde6dc5 100644 --- a/.gitignore +++ b/.gitignore @@ -128,6 +128,7 @@ deploy/docker-compose.override.yml .gocache/ vite.config.js docs/* +!docs/ACCOUNT_SCHEDULING_FLOW.md .serena/ # =================== diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 4e330cbf3a..07989265ce 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.101.3 +0.1.101.4 diff --git a/backend/internal/handler/admin/account_affinity_handler.go b/backend/internal/handler/admin/account_affinity_handler.go new file mode 100644 index 0000000000..62f4c3fe80 --- /dev/null +++ b/backend/internal/handler/admin/account_affinity_handler.go @@ -0,0 +1,254 @@ +package admin + +import ( + "context" + "log/slog" + "strconv" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/response" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +// 账号亲和调度 API 与契约说明见 docs/ACCOUNT_SCHEDULING_FLOW.md 。 + +// GetAffinityClients returns the list of affinity clients for an account with last active timestamps. +// GET /api/v1/admin/accounts/:id/affinity-clients +func (h *AccountHandler) GetAffinityClients(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + + account, err := h.adminService.GetAccount(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + + if !account.IsAffinityEnabled() { + response.Success(c, []service.AffinityClient{}) + return + } + + if h.gatewayCache == nil || len(account.GroupIDs) == 0 { + response.Success(c, []service.AffinityClient{}) + return + } + + clients, err := h.gatewayCache.GetAccountAffinityClientsWithScores( + c.Request.Context(), accountID, account.GroupIDs, service.ClientAffinityTTL, + ) + if err != nil { + response.Success(c, []service.AffinityClient{}) + return + } + + response.Success(c, clients) +} + +// clearAccountAffinity 清除指定账号在所有分组的亲和记录。 +func (h *AccountHandler) clearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) { + if h.gatewayCache == nil || len(groupIDs) == 0 { + return + } + if err := h.gatewayCache.ClearAccountAffinity(ctx, accountID, groupIDs); err != nil { + // 清理失败不影响主流程,记录日志即可 + slog.Warn("clear account affinity failed", + "account_id", accountID, + "error", err, + ) + } +} + +// countUniqueUsersFromAffinityMembers 从 "{userID}/{clientID}" 格式的成员列表中计算唯一用户数。 +func countUniqueUsersFromAffinityMembers(members []string) int64 { + users := make(map[string]struct{}, len(members)) + for _, m := range members { + if idx := strings.Index(m, "/"); idx > 0 { + users[m[:idx]] = struct{}{} + } + } + return int64(len(users)) +} + +// AffinityDetailClient 亲和详情中的单个客户端信息 +type AffinityDetailClient struct { + ClientID string `json:"client_id"` + LastActive time.Time `json:"last_active"` + IsPinned bool `json:"is_pinned"` +} + +// AffinityDetailUser 亲和详情中按用户分组的信息 +type AffinityDetailUser struct { + UserID int64 `json:"user_id"` + UserEmail string `json:"user_email"` + ClientCount int `json:"client_count"` + IsPinned bool `json:"is_pinned"` + Clients []AffinityDetailClient `json:"clients"` +} + +// AffinityDetailsResponse 亲和详情响应 +type AffinityDetailsResponse struct { + Users []AffinityDetailUser `json:"users"` + TotalUsers int `json:"total_users"` + TotalClients int `json:"total_clients"` + PinnedUsers []int64 `json:"pinned_users"` +} + +type affinityStateSnapshot struct { + enabled bool + groupIDs []int64 +} + +func (h *AccountHandler) captureAffinityStates(ctx context.Context, accountIDs []int64) map[int64]affinityStateSnapshot { + states := make(map[int64]affinityStateSnapshot, len(accountIDs)) + if h.gatewayCache == nil || len(accountIDs) == 0 { + return states + } + for _, accountID := range accountIDs { + account, err := h.adminService.GetAccount(ctx, accountID) + if err != nil || account == nil { + continue + } + states[accountID] = affinityStateSnapshot{ + enabled: account.IsAffinityEnabled(), + groupIDs: append([]int64(nil), account.GroupIDs...), + } + } + return states +} + +func (h *AccountHandler) clearAffinityCacheIfDisabled(ctx context.Context, accountID int64, oldState affinityStateSnapshot) { + if h.gatewayCache == nil || !oldState.enabled { + return + } + account, err := h.adminService.GetAccount(ctx, accountID) + if err != nil || account == nil || account.IsAffinityEnabled() { + return + } + groupIDs := oldState.groupIDs + if len(account.GroupIDs) > 0 { + groupIDs = mergeGroupIDs(groupIDs, account.GroupIDs) + } + h.clearAccountAffinity(ctx, accountID, groupIDs) +} + +func (h *AccountHandler) clearAffinityCacheForBulkIfDisabled(ctx context.Context, accountIDs []int64, oldStates map[int64]affinityStateSnapshot) { + for _, accountID := range accountIDs { + oldState, ok := oldStates[accountID] + if !ok { + continue + } + h.clearAffinityCacheIfDisabled(ctx, accountID, oldState) + } +} + +// GetAffinityDetails returns the affinity details grouped by user for an account. +// GET /api/v1/admin/accounts/:id/affinity-details +func (h *AccountHandler) GetAffinityDetails(c *gin.Context) { + accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) + if err != nil { + response.BadRequest(c, "Invalid account ID") + return + } + + emptyResp := AffinityDetailsResponse{ + Users: []AffinityDetailUser{}, + TotalUsers: 0, + TotalClients: 0, + PinnedUsers: []int64{}, + } + + account, err := h.adminService.GetAccount(c.Request.Context(), accountID) + if err != nil { + response.ErrorFrom(c, err) + return + } + + if !account.IsAffinityEnabled() { + response.Success(c, emptyResp) + return + } + + pinnedUsers := account.GetPinnedUsers() + if pinnedUsers == nil { + pinnedUsers = []int64{} + } + + if h.gatewayCache == nil || len(account.GroupIDs) == 0 { + emptyResp.PinnedUsers = pinnedUsers + response.Success(c, emptyResp) + return + } + + clients, err := h.gatewayCache.GetAccountAffinityClientsWithScores( + c.Request.Context(), accountID, account.GroupIDs, service.ClientAffinityTTL, + ) + if err != nil { + emptyResp.PinnedUsers = pinnedUsers + response.Success(c, emptyResp) + return + } + + pinnedSet := make(map[int64]struct{}, len(pinnedUsers)) + for _, uid := range pinnedUsers { + pinnedSet[uid] = struct{}{} + } + + // 按 UserID 分组 + userMap := make(map[int64]*AffinityDetailUser) + var userOrder []int64 + for _, cl := range clients { + u, ok := userMap[cl.UserID] + if !ok { + _, pinned := pinnedSet[cl.UserID] + u = &AffinityDetailUser{ + UserID: cl.UserID, + UserEmail: "", + ClientCount: 0, + IsPinned: pinned, + Clients: []AffinityDetailClient{}, + } + userMap[cl.UserID] = u + userOrder = append(userOrder, cl.UserID) + } + u.Clients = append(u.Clients, AffinityDetailClient{ + ClientID: cl.ClientID, + LastActive: cl.LastActive, + IsPinned: false, // 客户端级别暂无 pinned 概念 + }) + } + + // 关联用户邮箱(查询失败时保持空字符串,不影响主流程) + for _, uid := range userOrder { + user, uErr := h.adminService.GetUser(c.Request.Context(), uid) + if uErr != nil || user == nil { + continue + } + if u, ok := userMap[uid]; ok { + if user.Email != "" { + u.UserEmail = user.Email + } else { + u.UserEmail = user.Username + } + } + } + + users := make([]AffinityDetailUser, 0, len(userOrder)) + for _, uid := range userOrder { + u := userMap[uid] + u.ClientCount = len(u.Clients) + users = append(users, *u) + } + + response.Success(c, AffinityDetailsResponse{ + Users: users, + TotalUsers: len(users), + TotalClients: len(clients), + PinnedUsers: pinnedUsers, + }) +} diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 214afc5a9a..ce48977014 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -9,7 +9,6 @@ import ( "errors" "fmt" "log" - "log/slog" "net/http" "strconv" "strings" @@ -211,7 +210,7 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac } // 亲和客户端数据(启用亲和的账号始终返回 count,即使为 0) - if account.IsClientAffinityEnabled() { + if account.IsAffinityEnabled() { if h.gatewayCache != nil && len(account.GroupIDs) > 0 { accountGroups := map[int64][]int64{account.ID: account.GroupIDs} if clients, err := h.gatewayCache.GetAccountAffinityClientsBatch(ctx, accountGroups, service.ClientAffinityTTL); err == nil { @@ -219,17 +218,22 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac count := int64(len(cl)) item.AffinityClientCount = &count item.AffinityClients = cl + userCount := countUniqueUsersFromAffinityMembers(cl) + item.AffinityUserCount = &userCount } else { zero := int64(0) item.AffinityClientCount = &zero + item.AffinityUserCount = &zero } } else { zero := int64(0) item.AffinityClientCount = &zero + item.AffinityUserCount = &zero } } else { zero := int64(0) item.AffinityClientCount = &zero + item.AffinityUserCount = &zero } } @@ -351,7 +355,7 @@ func (h *AccountHandler) List(c *gin.Context) { accountGroups := make(map[int64][]int64) for i := range accounts { acc := &accounts[i] - if acc.IsClientAffinityEnabled() && len(acc.GroupIDs) > 0 { + if acc.IsAffinityEnabled() && len(acc.GroupIDs) > 0 { accountGroups[acc.ID] = acc.GroupIDs } } @@ -391,14 +395,18 @@ func (h *AccountHandler) List(c *gin.Context) { } // 注入亲和客户端数据到 DTO(启用亲和的账号始终返回 count,即使为 0) - if acc.IsClientAffinityEnabled() { + if acc.IsAffinityEnabled() { if clients, ok := affinityClients[acc.ID]; ok && len(clients) > 0 { count := int64(len(clients)) item.AffinityClientCount = &count item.AffinityClients = clients + // 从成员列表中解析唯一用户数 + userCount := countUniqueUsersFromAffinityMembers(clients) + item.AffinityUserCount = &userCount } else { zero := int64(0) item.AffinityClientCount = &zero + item.AffinityUserCount = &zero } } @@ -622,15 +630,7 @@ func (h *AccountHandler) Update(c *gin.Context) { // base_rpm 输入校验:负值归零,超过 10000 截断 sanitizeExtraBaseRPM(req.Extra) - // 记录更新前的亲和状态,用于检测亲和关闭时清理 Redis 记录 - oldAffinityEnabled := false - var oldGroupIDs []int64 - if len(req.Extra) > 0 && h.gatewayCache != nil { - if oldAccount, err := h.adminService.GetAccount(c.Request.Context(), accountID); err == nil { - oldAffinityEnabled = oldAccount.IsClientAffinityEnabled() - oldGroupIDs = oldAccount.GroupIDs - } - } + oldStates := h.captureAffinityStates(c.Request.Context(), []int64{accountID}) // 确定是否跳过混合渠道检查 skipCheck := req.ConfirmMixedChannelRisk != nil && *req.ConfirmMixedChannelRisk @@ -668,14 +668,7 @@ func (h *AccountHandler) Update(c *gin.Context) { return } - // 亲和关闭时清理 Redis 中的亲和记录 - if oldAffinityEnabled && !account.IsClientAffinityEnabled() { - groupIDs := oldGroupIDs - if len(account.GroupIDs) > 0 { - groupIDs = mergeGroupIDs(oldGroupIDs, account.GroupIDs) - } - h.clearAccountAffinity(c.Request.Context(), accountID, groupIDs) - } + h.clearAffinityCacheForBulkIfDisabled(c.Request.Context(), []int64{accountID}, oldStates) response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account)) } @@ -1393,6 +1386,8 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) { return } + oldStates := h.captureAffinityStates(c.Request.Context(), req.AccountIDs) + result, err := h.adminService.BulkUpdateAccounts(c.Request.Context(), &service.BulkUpdateAccountsInput{ AccountIDs: req.AccountIDs, Name: req.Name, @@ -1427,6 +1422,8 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) { return } + h.clearAffinityCacheForBulkIfDisabled(c.Request.Context(), req.AccountIDs, oldStates) + response.Success(c, result) } @@ -1639,56 +1636,6 @@ func (h *AccountHandler) ResetQuota(c *gin.Context) { response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account)) } -// GetAffinityClients returns the list of affinity clients for an account with last active timestamps. -// GET /api/v1/admin/accounts/:id/affinity-clients -func (h *AccountHandler) GetAffinityClients(c *gin.Context) { - accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) - if err != nil { - response.BadRequest(c, "Invalid account ID") - return - } - - account, err := h.adminService.GetAccount(c.Request.Context(), accountID) - if err != nil { - response.ErrorFrom(c, err) - return - } - - if !account.IsClientAffinityEnabled() { - response.Success(c, []service.AffinityClient{}) - return - } - - if h.gatewayCache == nil || len(account.GroupIDs) == 0 { - response.Success(c, []service.AffinityClient{}) - return - } - - clients, err := h.gatewayCache.GetAccountAffinityClientsWithScores( - c.Request.Context(), accountID, account.GroupIDs, service.ClientAffinityTTL, - ) - if err != nil { - response.Success(c, []service.AffinityClient{}) - return - } - - response.Success(c, clients) -} - -// clearAccountAffinity 清除指定账号在所有分组的亲和记录。 -func (h *AccountHandler) clearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) { - if h.gatewayCache == nil || len(groupIDs) == 0 { - return - } - if err := h.gatewayCache.ClearAccountAffinity(ctx, accountID, groupIDs); err != nil { - // 清理失败不影响主流程,记录日志即可 - slog.Warn("clear account affinity failed", - "account_id", accountID, - "error", err, - ) - } -} - // mergeGroupIDs 合并两个 groupID 切片并去重。 func mergeGroupIDs(a, b []int64) []int64 { seen := make(map[int64]struct{}, len(a)+len(b)) diff --git a/backend/internal/handler/admin/account_handler_affinity_cache_test.go b/backend/internal/handler/admin/account_handler_affinity_cache_test.go new file mode 100644 index 0000000000..3c22bb543a --- /dev/null +++ b/backend/internal/handler/admin/account_handler_affinity_cache_test.go @@ -0,0 +1,298 @@ +package admin + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type affinityCacheAdminService struct { + *stubAdminService + accounts map[int64]*service.Account + bulkResultIDs []int64 + bulkFailedCount int +} + +func (s *affinityCacheAdminService) GetAccount(_ context.Context, id int64) (*service.Account, error) { + if acc, ok := s.accounts[id]; ok { + cloned := *acc + return &cloned, nil + } + return s.stubAdminService.GetAccount(context.Background(), id) +} + +func (s *affinityCacheAdminService) UpdateAccount(_ context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) { + acc := s.accounts[id] + if acc == nil { + return s.stubAdminService.UpdateAccount(context.Background(), id, input) + } + if input.Extra != nil { + acc.Extra = input.Extra + } + if input.GroupIDs != nil { + acc.GroupIDs = *input.GroupIDs + } + cloned := *acc + return &cloned, nil +} + +func (s *affinityCacheAdminService) BulkUpdateAccounts(_ context.Context, input *service.BulkUpdateAccountsInput) (*service.BulkUpdateAccountsResult, error) { + for _, accountID := range input.AccountIDs { + acc := s.accounts[accountID] + if acc == nil { + continue + } + if input.Extra != nil { + acc.Extra = input.Extra + } + if input.GroupIDs != nil { + acc.GroupIDs = *input.GroupIDs + } + } + successIDs := append([]int64(nil), input.AccountIDs...) + failed := 0 + if len(s.bulkResultIDs) > 0 { + successIDs = append([]int64(nil), s.bulkResultIDs...) + failed = s.bulkFailedCount + } + return &service.BulkUpdateAccountsResult{ + Success: len(successIDs), + Failed: failed, + SuccessIDs: successIDs, + Results: []service.BulkUpdateAccountResult{}, + }, nil +} + +type affinityCacheClearRecorder struct { + clearCalls []struct { + accountID int64 + groupIDs []int64 + } +} + +func (r *affinityCacheClearRecorder) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) { + return 0, nil +} +func (r *affinityCacheClearRecorder) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (r *affinityCacheClearRecorder) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error { + return nil +} +func (r *affinityCacheClearRecorder) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { + return nil +} +func (r *affinityCacheClearRecorder) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (r *affinityCacheClearRecorder) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (r *affinityCacheClearRecorder) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (r *affinityCacheClearRecorder) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (r *affinityCacheClearRecorder) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]service.AffinityClient, error) { + return nil, nil +} +func (r *affinityCacheClearRecorder) ClearAccountAffinity(_ context.Context, accountID int64, groupIDs []int64) error { + r.clearCalls = append(r.clearCalls, struct { + accountID int64 + groupIDs []int64 + }{ + accountID: accountID, + groupIDs: append([]int64(nil), groupIDs...), + }) + return nil +} +func (r *affinityCacheClearRecorder) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (users, clients, perUser int64, err error) { + return 0, 0, 0, nil +} + +func setupAffinityCacheRouter(adminSvc service.AdminService, gatewayCache service.GatewayCache) *gin.Engine { + gin.SetMode(gin.TestMode) + router := gin.New() + handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, gatewayCache) + router.PUT("/api/v1/admin/accounts/:id", handler.Update) + router.POST("/api/v1/admin/accounts/bulk-update", handler.BulkUpdate) + return router +} + +func TestAccountHandlerUpdate_DisablingAffinityClearsCache(t *testing.T) { + adminSvc := &affinityCacheAdminService{ + stubAdminService: newStubAdminService(), + accounts: map[int64]*service.Account{ + 3: { + ID: 3, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{1, 2}, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + }, + } + cache := &affinityCacheClearRecorder{} + router := setupAffinityCacheRouter(adminSvc, cache) + + body, _ := json.Marshal(map[string]any{ + "extra": map[string]any{}, + }) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPut, "/api/v1/admin/accounts/3", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Len(t, cache.clearCalls, 1) + require.Equal(t, int64(3), cache.clearCalls[0].accountID) + require.Equal(t, []int64{1, 2}, cache.clearCalls[0].groupIDs) +} + +func TestAccountHandlerBulkUpdate_DisablingAffinityClearsCache(t *testing.T) { + adminSvc := &affinityCacheAdminService{ + stubAdminService: newStubAdminService(), + accounts: map[int64]*service.Account{ + 1: { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{11}, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + 2: { + ID: 2, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeSetupToken, + Status: service.StatusActive, + GroupIDs: []int64{22}, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + }, + } + cache := &affinityCacheClearRecorder{} + router := setupAffinityCacheRouter(adminSvc, cache) + + body, _ := json.Marshal(map[string]any{ + "account_ids": []int64{1, 2}, + "extra": map[string]any{ + "client_affinity_enabled": false, + }, + }) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Len(t, cache.clearCalls, 2) + require.Equal(t, int64(1), cache.clearCalls[0].accountID) + require.Equal(t, []int64{11}, cache.clearCalls[0].groupIDs) + require.Equal(t, int64(2), cache.clearCalls[1].accountID) + require.Equal(t, []int64{22}, cache.clearCalls[1].groupIDs) +} + +func TestAccountHandlerBulkUpdate_DisablingAffinityWithNewFlagClearsCache(t *testing.T) { + adminSvc := &affinityCacheAdminService{ + stubAdminService: newStubAdminService(), + accounts: map[int64]*service.Account{ + 1: { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{33}, + Extra: map[string]any{ + "affinity_enabled": true, + }, + }, + }, + } + cache := &affinityCacheClearRecorder{} + router := setupAffinityCacheRouter(adminSvc, cache) + + body, _ := json.Marshal(map[string]any{ + "account_ids": []int64{1}, + "extra": map[string]any{ + "affinity_enabled": false, + }, + }) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Len(t, cache.clearCalls, 1) + require.Equal(t, int64(1), cache.clearCalls[0].accountID) + require.Equal(t, []int64{33}, cache.clearCalls[0].groupIDs) +} + +func TestAccountHandlerBulkUpdate_PartialFailureStillClearsDisabledAffinityCache(t *testing.T) { + adminSvc := &affinityCacheAdminService{ + stubAdminService: newStubAdminService(), + accounts: map[int64]*service.Account{ + 1: { + ID: 1, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{44}, + Extra: map[string]any{ + "affinity_enabled": true, + }, + }, + 2: { + ID: 2, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeSetupToken, + Status: service.StatusActive, + GroupIDs: []int64{55}, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + }, + bulkResultIDs: []int64{1}, + bulkFailedCount: 1, + } + cache := &affinityCacheClearRecorder{} + router := setupAffinityCacheRouter(adminSvc, cache) + + body, _ := json.Marshal(map[string]any{ + "account_ids": []int64{1, 2}, + "group_ids": []int64{101}, + "extra": map[string]any{ + "affinity_enabled": false, + "client_affinity_enabled": false, + }, + }) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Len(t, cache.clearCalls, 2) + require.Equal(t, int64(1), cache.clearCalls[0].accountID) + require.Equal(t, []int64{44, 101}, cache.clearCalls[0].groupIDs) + require.Equal(t, int64(2), cache.clearCalls[1].accountID) + require.Equal(t, []int64{55, 101}, cache.clearCalls[1].groupIDs) +} diff --git a/backend/internal/handler/admin/account_handler_affinity_details_test.go b/backend/internal/handler/admin/account_handler_affinity_details_test.go new file mode 100644 index 0000000000..0a4b52d579 --- /dev/null +++ b/backend/internal/handler/admin/account_handler_affinity_details_test.go @@ -0,0 +1,205 @@ +package admin + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type affinityDetailsAdminService struct { + *stubAdminService + account service.Account + usersByID map[int64]service.User +} + +func (s *affinityDetailsAdminService) GetAccount(_ context.Context, id int64) (*service.Account, error) { + if s.account.ID == id { + acc := s.account + return &acc, nil + } + return s.stubAdminService.GetAccount(context.Background(), id) +} + +func (s *affinityDetailsAdminService) GetUser(_ context.Context, id int64) (*service.User, error) { + if u, ok := s.usersByID[id]; ok { + user := u + return &user, nil + } + return s.stubAdminService.GetUser(context.Background(), id) +} + +type affinityDetailsGatewayCacheStub struct { + clients []service.AffinityClient +} + +func (s *affinityDetailsGatewayCacheStub) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) { + return 0, nil +} +func (s *affinityDetailsGatewayCacheStub) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil +} +func (s *affinityDetailsGatewayCacheStub) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { + return map[int64]int64{}, nil +} +func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) { + return map[int64][]string{}, nil +} +func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]service.AffinityClient, error) { + return s.clients, nil +} +func (s *affinityDetailsGatewayCacheStub) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} +func (s *affinityDetailsGatewayCacheStub) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (users, clients, perUser int64, err error) { + return 0, 0, 0, nil +} + +func setupAffinityDetailsRouter(adminSvc service.AdminService, gatewayCache service.GatewayCache) *gin.Engine { + gin.SetMode(gin.TestMode) + router := gin.New() + handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, gatewayCache) + router.GET("/api/v1/admin/accounts/:id/affinity-details", handler.GetAffinityDetails) + return router +} + +func TestAccountHandlerGetAffinityDetails_ContractFields(t *testing.T) { + adminSvc := &affinityDetailsAdminService{ + stubAdminService: newStubAdminService(), + account: service.Account{ + ID: 88, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{1}, + Extra: map[string]any{ + "client_affinity_enabled": true, + "pinned_users": []any{float64(42), float64(99)}, + }, + }, + usersByID: map[int64]service.User{ + 42: {ID: 42, Email: "user42@example.com"}, + 7: {ID: 7, Username: "user7"}, + }, + } + cache := &affinityDetailsGatewayCacheStub{ + clients: []service.AffinityClient{ + {UserID: 42, ClientID: "client-a", LastActive: time.Now().Add(-time.Minute)}, + {UserID: 7, ClientID: "client-b", LastActive: time.Now().Add(-2 * time.Minute)}, + {UserID: 42, ClientID: "client-c", LastActive: time.Now().Add(-3 * time.Minute)}, + }, + } + router := setupAffinityDetailsRouter(adminSvc, cache) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/88/affinity-details", nil) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var resp struct { + Code int `json:"code"` + Data struct { + Users []struct { + UserID int64 `json:"user_id"` + UserEmail string `json:"user_email"` + ClientCount int `json:"client_count"` + IsPinned bool `json:"is_pinned"` + Clients []struct { + ClientID string `json:"client_id"` + } `json:"clients"` + } `json:"users"` + TotalUsers int `json:"total_users"` + TotalClients int `json:"total_clients"` + PinnedUsers []int64 `json:"pinned_users"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, 0, resp.Code) + require.Equal(t, 2, resp.Data.TotalUsers) + require.Equal(t, 3, resp.Data.TotalClients) + require.Equal(t, []int64{42, 99}, resp.Data.PinnedUsers) + require.Len(t, resp.Data.Users, 2) + + usersByID := make(map[int64]struct { + UserEmail string + ClientCount int + IsPinned bool + ClientLen int + }, len(resp.Data.Users)) + for _, u := range resp.Data.Users { + usersByID[u.UserID] = struct { + UserEmail string + ClientCount int + IsPinned bool + ClientLen int + }{ + UserEmail: u.UserEmail, + ClientCount: u.ClientCount, + IsPinned: u.IsPinned, + ClientLen: len(u.Clients), + } + } + require.Equal(t, "user42@example.com", usersByID[42].UserEmail) + require.Equal(t, 2, usersByID[42].ClientCount) + require.True(t, usersByID[42].IsPinned) + require.Equal(t, 2, usersByID[42].ClientLen) + + require.Equal(t, "user7", usersByID[7].UserEmail) + require.Equal(t, 1, usersByID[7].ClientCount) + require.False(t, usersByID[7].IsPinned) + require.Equal(t, 1, usersByID[7].ClientLen) +} + +func TestAccountHandlerGetAffinityDetails_DisabledReturnsEmptyContract(t *testing.T) { + adminSvc := &affinityDetailsAdminService{ + stubAdminService: newStubAdminService(), + account: service.Account{ + ID: 89, + Platform: service.PlatformAnthropic, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + GroupIDs: []int64{1}, + }, + usersByID: map[int64]service.User{}, + } + router := setupAffinityDetailsRouter(adminSvc, &affinityDetailsGatewayCacheStub{}) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/89/affinity-details", nil) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var resp struct { + Code int `json:"code"` + Data struct { + Users []any `json:"users"` + TotalUsers int `json:"total_users"` + TotalClients int `json:"total_clients"` + PinnedUsers []int64 `json:"pinned_users"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, 0, resp.Code) + require.Empty(t, resp.Data.Users) + require.Equal(t, 0, resp.Data.TotalUsers) + require.Equal(t, 0, resp.Data.TotalClients) + require.Empty(t, resp.Data.PinnedUsers) +} diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 447d909558..fd9b9d66dd 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -266,9 +266,14 @@ func AccountFromServiceShallow(a *service.Account) *Account { } // 客户端亲和调度(Anthropic 和 Antigravity 账号) - if a.IsClientAffinityEnabled() { + if a.IsAffinityEnabled() { enabled := true out.ClientAffinityEnabled = &enabled + allow := a.IsAffinityAllowSwitch() + out.AffinityAllowSwitch = &allow + if pinnedUsers := a.GetPinnedUsers(); len(pinnedUsers) > 0 { + out.PinnedUserIDs = pinnedUsers + } } // 提取账号配额限制(apikey / bedrock 类型有效) diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index db1e9a610c..01ec2eba76 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -201,6 +201,15 @@ type Account struct { // 启用后新会话会优先调度到客户端之前使用过的账号 ClientAffinityEnabled *bool `json:"client_affinity_enabled,omitempty"` + // 亲和允许切换(默认 true) + AffinityAllowSwitch *bool `json:"affinity_allow_switch,omitempty"` + + // 亲和用户数量(admin 列表端点注入) + AffinityUserCount *int64 `json:"affinity_user_count,omitempty"` + + // 指定亲和用户 ID 列表 + PinnedUserIDs []int64 `json:"pinned_user_ids,omitempty"` + // 亲和客户端数据(仅 admin 列表端点注入,不由 mapper 填充) AffinityClientCount *int64 `json:"affinity_client_count,omitempty"` AffinityClients []string `json:"affinity_clients,omitempty"` diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index cb90f49ba9..80d9d58dd4 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -291,7 +291,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { } for { - selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, "") // Gemini 不使用会话限制 + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, "", int64(0)) // Gemini 不使用会话限制 if err != nil { if len(fs.FailedAccountIDs) == 0 { h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts: "+err.Error(), streamStarted) @@ -501,8 +501,16 @@ func (h *GatewayHandler) Messages(c *gin.Context) { for { // 选择支持该模型的账号 - selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), currentAPIKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, parsedReq.MetadataUserID) + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), currentAPIKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, parsedReq.MetadataUserID, subject.UserID) if err != nil { + if errors.Is(err, service.ErrAffinityNoSwitch) { + h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "Affinity account unavailable and switching is disabled", streamStarted) + return + } + if errors.Is(err, service.ErrAffinityLimitExceeded) { + h.handleStreamingAwareError(c, http.StatusTooManyRequests, "api_error", "Affinity client limit exceeded", streamStarted) + return + } if len(fs.FailedAccountIDs) == 0 { h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts: "+err.Error(), streamStarted) return diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index cfe809114b..4dc3b282b5 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -352,7 +352,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { } for { - selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, modelName, fs.FailedAccountIDs, "") // Gemini 不使用会话限制 + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, modelName, fs.FailedAccountIDs, "", int64(0)) // Gemini 不使用会话限制 if err != nil { if len(fs.FailedAccountIDs) == 0 { googleError(c, http.StatusServiceUnavailable, "No available Gemini accounts: "+err.Error()) diff --git a/backend/internal/handler/sora_gateway_handler.go b/backend/internal/handler/sora_gateway_handler.go index dc301ce149..3cc6a0397e 100644 --- a/backend/internal/handler/sora_gateway_handler.go +++ b/backend/internal/handler/sora_gateway_handler.go @@ -225,7 +225,7 @@ func (h *SoraGatewayHandler) ChatCompletions(c *gin.Context) { var lastFailoverHeaders http.Header for { - selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionHash, reqModel, failedAccountIDs, "") + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionHash, reqModel, failedAccountIDs, "", int64(0)) if err != nil { reqLog.Warn("sora.account_select_failed", zap.Error(err), diff --git a/backend/internal/handler/sora_videos_handler.go b/backend/internal/handler/sora_videos_handler.go index 46edbd3624..b2a9d79c81 100644 --- a/backend/internal/handler/sora_videos_handler.go +++ b/backend/internal/handler/sora_videos_handler.go @@ -323,7 +323,7 @@ func (h *SoraVideosHandler) selectAccount(c *gin.Context, model string) (*servic } selection, err := h.gatewayService.SelectAccountWithLoadAwareness( - c.Request.Context(), apiKey.GroupID, "", model, nil, "", + c.Request.Context(), apiKey.GroupID, "", model, nil, "", int64(0), ) if err != nil { soraErrorResponse(c, http.StatusServiceUnavailable, "server_error", "No available accounts") diff --git a/backend/internal/repository/gateway_cache.go b/backend/internal/repository/gateway_cache.go index ec4bf40e3e..e40baffdc0 100644 --- a/backend/internal/repository/gateway_cache.go +++ b/backend/internal/repository/gateway_cache.go @@ -5,6 +5,7 @@ import ( _ "embed" "fmt" "strconv" + "strings" "time" "github.com/Wei-Shaw/sub2api/internal/service" @@ -12,9 +13,9 @@ import ( ) const ( - stickySessionPrefix = "sticky_session:" - clientAffinityPrefix = "client_affinity:" - clientAffinityReversePrefix = "client_affinity_rev:" + stickySessionPrefix = "sticky_session:" + affinityKeyPrefix = "affinity:" + affinityRevKeyPrefix = "affinity_rev:" ) var ( @@ -30,6 +31,8 @@ var ( getAffinityClientsWithScoresLua string //go:embed lua/clear_account_affinity.lua clearAccountAffinityLua string + //go:embed lua/get_affinity_multi_count.lua + getAffinityMultiCountLua string getAffinityScript = redis.NewScript(getAffinityLua) updateAffinityScript = redis.NewScript(updateAffinityLua) @@ -37,6 +40,7 @@ var ( getAffinityClientsScript = redis.NewScript(getAffinityClientsLua) getAffinityClientsWithScoresScript = redis.NewScript(getAffinityClientsWithScoresLua) clearAccountAffinityScript = redis.NewScript(clearAccountAffinityLua) + getAffinityMultiCountScript = redis.NewScript(getAffinityMultiCountLua) ) type gatewayCache struct { @@ -84,20 +88,39 @@ func (c *gatewayCache) DeleteSessionAccountID(ctx context.Context, groupID int64 return c.rdb.Del(ctx, key).Err() } -// buildAffinityKey 构建正向亲和 key(client → accounts) -// 格式: client_affinity:{groupID}:{clientID} -func buildAffinityKey(groupID int64, clientID string) string { - return fmt.Sprintf("%s%d:%s", clientAffinityPrefix, groupID, clientID) +// buildAffinityKey 构建正向亲和 key(member → accounts) +// 格式: affinity:{groupID}:{userID}/{clientID} +func buildAffinityKey(groupID int64, userID int64, clientID string) string { + return fmt.Sprintf("%s%d:%s", affinityKeyPrefix, groupID, buildAffinityMember(userID, clientID)) } -// buildAffinityReverseKey 构建反向亲和 key(account → clients) -// 格式: client_affinity_rev:{groupID}:{accountID} +// buildAffinityReverseKey 构建反向亲和 key(account → members) +// 格式: affinity_rev:{groupID}:{accountID} func buildAffinityReverseKey(groupID int64, accountID int64) string { - return fmt.Sprintf("%s%d:%d", clientAffinityReversePrefix, groupID, accountID) + return fmt.Sprintf("%s%d:%d", affinityRevKeyPrefix, groupID, accountID) } -func (c *gatewayCache) GetClientAffinityAccounts(ctx context.Context, groupID int64, clientID string, ttl time.Duration) ([]int64, error) { - key := buildAffinityKey(groupID, clientID) +// buildAffinityMember 构建亲和成员标识 +// 格式: {userID}/{clientID} +func buildAffinityMember(userID int64, clientID string) string { + return fmt.Sprintf("%d/%s", userID, clientID) +} + +// parseAffinityMember 解析亲和成员标识为 userID 和 clientID +func parseAffinityMember(member string) (userID int64, clientID string) { + idx := strings.IndexByte(member, '/') + if idx < 0 { + // 兼容旧格式(纯 clientID,无 userID 前缀) + return 0, member + } + userID, _ = strconv.ParseInt(member[:idx], 10, 64) + clientID = member[idx+1:] + return userID, clientID +} + +// GetAffinityAccounts 获取亲和账号列表(按最近使用降序),同时清理过期成员 +func (c *gatewayCache) GetAffinityAccounts(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) { + key := buildAffinityKey(groupID, userID, clientID) now := time.Now().Unix() expireThreshold := now - int64(ttl.Seconds()) @@ -120,19 +143,21 @@ func (c *gatewayCache) GetClientAffinityAccounts(ctx context.Context, groupID in return accountIDs, nil } -func (c *gatewayCache) UpdateClientAffinity(ctx context.Context, groupID int64, clientID string, accountID int64, ttl time.Duration) error { - fwdKey := buildAffinityKey(groupID, clientID) +// UpdateAffinity 添加/更新亲和关系(更新 score 为当前时间戳,刷新 key TTL) +func (c *gatewayCache) UpdateAffinity(ctx context.Context, groupID int64, userID int64, clientID string, accountID int64, ttl time.Duration) error { + fwdKey := buildAffinityKey(groupID, userID, clientID) revKey := buildAffinityReverseKey(groupID, accountID) now := time.Now().Unix() ttlSeconds := int64(ttl.Seconds()) expireThreshold := now - ttlSeconds + member := buildAffinityMember(userID, clientID) return updateAffinityScript.Run(ctx, c.rdb, []string{fwdKey, revKey}, - now, ttlSeconds, accountID, expireThreshold, clientID, + now, ttlSeconds, accountID, expireThreshold, member, ).Err() } -// GetAccountAffinityCountBatch 批量获取账号的亲和客户端数量(惰性清理过期成员) +// GetAccountAffinityCountBatch 批量获取账号的亲和成员数量(惰性清理过期成员) func (c *gatewayCache) GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) { if len(accountIDs) == 0 { return map[int64]int64{}, nil @@ -162,8 +187,9 @@ func (c *gatewayCache) GetAccountAffinityCountBatch(ctx context.Context, groupID return result, nil } -// GetAccountAffinityClientsBatch 批量获取每个账号跨所有分组的亲和客户端列表(去重)。 +// GetAccountAffinityClientsBatch 批量获取每个账号跨所有分组的亲和成员列表(去重)。 // accountGroups: map[accountID][]groupID,对每个 (groupID, accountID) 组合查询反向索引。 +// 返回值成员格式为 {userID}/{clientID}。 func (c *gatewayCache) GetAccountAffinityClientsBatch(ctx context.Context, accountGroups map[int64][]int64, ttl time.Duration) (map[int64][]string, error) { if len(accountGroups) == 0 { return map[int64][]string{}, nil @@ -197,21 +223,21 @@ func (c *gatewayCache) GetAccountAffinityClientsBatch(ctx context.Context, accou return nil, err } - // 合并结果:同一个 accountID 跨多个 group 的 clientID 去重 + // 合并结果:同一个 accountID 跨多个 group 的成员去重 result := make(map[int64][]string, len(accountGroups)) seen := make(map[int64]map[string]struct{}, len(accountGroups)) for i, q := range queries { - clients, _ := cmds[i].StringSlice() - if len(clients) == 0 { + members, _ := cmds[i].StringSlice() + if len(members) == 0 { continue } if seen[q.accountID] == nil { seen[q.accountID] = make(map[string]struct{}) } - for _, clientID := range clients { - if _, exists := seen[q.accountID][clientID]; !exists { - seen[q.accountID][clientID] = struct{}{} - result[q.accountID] = append(result[q.accountID], clientID) + for _, member := range members { + if _, exists := seen[q.accountID][member]; !exists { + seen[q.accountID][member] = struct{}{} + result[q.accountID] = append(result[q.accountID], member) } } } @@ -245,25 +271,32 @@ func (c *gatewayCache) GetAccountAffinityClientsWithScores( return nil, err } - // 合并跨组结果,同一 clientID 取最近的 lastActive - seen := make(map[string]int64) // clientID → max timestamp + // 合并跨组结果,同一 member 取最近的 lastActive + type memberInfo struct { + userID int64 + clientID string + ts int64 + } + seen := make(map[string]*memberInfo) // member string → info for _, cmd := range cmds { vals, _ := cmd.StringSlice() - // vals 格式: [clientID1, score1, clientID2, score2, ...] + // vals 格式: [member1, score1, member2, score2, ...] for j := 0; j+1 < len(vals); j += 2 { - clientID := vals[j] + member := vals[j] ts, _ := strconv.ParseInt(vals[j+1], 10, 64) - if existing, ok := seen[clientID]; !ok || ts > existing { - seen[clientID] = ts + if existing, ok := seen[member]; !ok || ts > existing.ts { + uid, cid := parseAffinityMember(member) + seen[member] = &memberInfo{userID: uid, clientID: cid, ts: ts} } } } result := make([]service.AffinityClient, 0, len(seen)) - for clientID, ts := range seen { + for _, info := range seen { result = append(result, service.AffinityClient{ - ClientID: clientID, - LastActive: time.Unix(ts, 0), + UserID: info.userID, + ClientID: info.clientID, + LastActive: time.Unix(info.ts, 0), }) } @@ -274,8 +307,8 @@ func (c *gatewayCache) GetAccountAffinityClientsWithScores( } // ClearAccountAffinity 清除指定账号在所有分组的亲和记录(正向+反向索引)。 -// 对每个 groupID 执行 Lua 脚本:读取反向索引获取所有客户端, -// 从每个客户端的正向索引中移除该账号,然后删除反向索引。 +// 对每个 groupID 执行 Lua 脚本:读取反向索引获取所有成员, +// 从每个成员的正向索引中移除该账号,然后删除反向索引。 func (c *gatewayCache) ClearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) error { if len(groupIDs) == 0 { return nil @@ -294,3 +327,37 @@ func (c *gatewayCache) ClearAccountAffinity(ctx context.Context, accountID int64 } return nil } + +// GetAffinityMultiCount 获取账号的多维度亲和计数。 +// 返回: uniqueUsers(独立用户数), uniqueClients(独立客户端数), perUserClients(目标用户的客户端数) +func (c *gatewayCache) GetAffinityMultiCount( + ctx context.Context, + groupID int64, + accountID int64, + targetUserID int64, + ttl time.Duration, +) (users, clients, perUser int64, err error) { + key := buildAffinityReverseKey(groupID, accountID) + now := time.Now().Unix() + expireThreshold := now - int64(ttl.Seconds()) + + targetUserStr := "" + if targetUserID > 0 { + targetUserStr = strconv.FormatInt(targetUserID, 10) + } + + result, err := getAffinityMultiCountScript.Run(ctx, c.rdb, []string{key}, expireThreshold, targetUserStr).Int64Slice() + if err != nil { + if err == redis.Nil { + return 0, 0, 0, nil + } + return 0, 0, 0, err + } + + if len(result) < 4 { + return 0, 0, 0, nil + } + + // result: {totalMembers, uniqueUsers, uniqueClients, perUserClients} + return result[1], result[2], result[3], nil +} diff --git a/backend/internal/repository/lua/clear_account_affinity.lua b/backend/internal/repository/lua/clear_account_affinity.lua index e125be1690..d2f0af0338 100644 --- a/backend/internal/repository/lua/clear_account_affinity.lua +++ b/backend/internal/repository/lua/clear_account_affinity.lua @@ -1,21 +1,21 @@ -- 清除单个账号在指定分组的所有亲和记录(正向+反向) --- KEYS[1] = client_affinity_rev:{groupID}:{accountID} (反向索引) +-- KEYS[1] = affinity_rev:{groupID}:{accountID} (反向索引) -- ARGV[1] = groupID (用于构建正向 key) -- ARGV[2] = accountID (正向索引中要移除的成员) --- 返回: 清理的客户端数量 +-- 返回: 清理的成员数量 local rev_key = KEYS[1] local group_id = ARGV[1] local account_id = ARGV[2] --- 获取反向索引中所有客户端 ID -local clients = redis.call('ZRANGE', rev_key, 0, -1) -if #clients == 0 then +-- 获取反向索引中所有成员 ({userID}/{clientID}) +local members = redis.call('ZRANGE', rev_key, 0, -1) +if #members == 0 then return 0 end --- 从每个客户端的正向索引中移除该账号 -for _, client_id in ipairs(clients) do - local fwd_key = 'client_affinity:' .. group_id .. ':' .. client_id +-- 从每个成员的正向索引中移除该账号 +for _, member in ipairs(members) do + local fwd_key = 'affinity:' .. group_id .. ':' .. member redis.call('ZREM', fwd_key, account_id) -- 如果正向索引为空,删除 key if redis.call('ZCARD', fwd_key) == 0 then @@ -26,4 +26,4 @@ end -- 删除反向索引 redis.call('DEL', rev_key) -return #clients +return #members diff --git a/backend/internal/repository/lua/get_affinity.lua b/backend/internal/repository/lua/get_affinity.lua index 7e3971dbea..9db85be2a9 100644 --- a/backend/internal/repository/lua/get_affinity.lua +++ b/backend/internal/repository/lua/get_affinity.lua @@ -1,5 +1,5 @@ -- 清理过期成员后返回亲和账号列表(按最近使用降序) --- KEYS[1] = client_affinity:{groupID}:{clientID} +-- KEYS[1] = affinity:{groupID}:{userID}/{clientID} -- ARGV[1] = 过期阈值时间戳 (now - ttl) redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) return redis.call('ZREVRANGE', KEYS[1], 0, -1) diff --git a/backend/internal/repository/lua/get_affinity_clients.lua b/backend/internal/repository/lua/get_affinity_clients.lua index 049b1c5b7c..7290b63758 100644 --- a/backend/internal/repository/lua/get_affinity_clients.lua +++ b/backend/internal/repository/lua/get_affinity_clients.lua @@ -1,5 +1,6 @@ --- 清理过期成员后返回反向索引的 clientID 列表(按最近使用降序) --- KEYS[1] = client_affinity_rev:{groupID}:{accountID} +-- 清理过期成员后返回反向索引的成员列表(按最近使用降序) +-- 成员格式: {userID}/{clientID} +-- KEYS[1] = affinity_rev:{groupID}:{accountID} -- ARGV[1] = 过期阈值时间戳 (now - ttl) redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) return redis.call('ZREVRANGE', KEYS[1], 0, -1) diff --git a/backend/internal/repository/lua/get_affinity_clients_with_scores.lua b/backend/internal/repository/lua/get_affinity_clients_with_scores.lua index 8e70708efb..bff6b616ed 100644 --- a/backend/internal/repository/lua/get_affinity_clients_with_scores.lua +++ b/backend/internal/repository/lua/get_affinity_clients_with_scores.lua @@ -1,6 +1,7 @@ --- 清理过期成员后返回反向索引的 clientID 列表及其 score(最后活跃时间戳) --- KEYS[1] = client_affinity_rev:{groupID}:{accountID} +-- 清理过期成员后返回反向索引的成员列表及其 score(最后活跃时间戳) +-- 成员格式: {userID}/{clientID} +-- KEYS[1] = affinity_rev:{groupID}:{accountID} -- ARGV[1] = 过期阈值时间戳 (now - ttl) --- 返回: {clientID1, score1, clientID2, score2, ...}(按最近使用降序) +-- 返回: {member1, score1, member2, score2, ...}(按最近使用降序) redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) return redis.call('ZREVRANGEBYSCORE', KEYS[1], '+inf', '-inf', 'WITHSCORES') diff --git a/backend/internal/repository/lua/get_affinity_count.lua b/backend/internal/repository/lua/get_affinity_count.lua index 4b00c22daa..7cf9b5cb51 100644 --- a/backend/internal/repository/lua/get_affinity_count.lua +++ b/backend/internal/repository/lua/get_affinity_count.lua @@ -1,5 +1,5 @@ -- 清理过期成员后返回反向索引的成员数量 --- KEYS[1] = client_affinity_rev:{groupID}:{accountID} +-- KEYS[1] = affinity_rev:{groupID}:{accountID} -- ARGV[1] = 过期阈值时间戳 (now - ttl) redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) return redis.call('ZCARD', KEYS[1]) diff --git a/backend/internal/repository/lua/get_affinity_multi_count.lua b/backend/internal/repository/lua/get_affinity_multi_count.lua new file mode 100644 index 0000000000..166cf4f9a3 --- /dev/null +++ b/backend/internal/repository/lua/get_affinity_multi_count.lua @@ -0,0 +1,40 @@ +-- 从反向索引解析多维度计数(用户/客户端/每用户客户端) +-- KEYS[1] = affinity_rev:{groupID}:{accountID} +-- ARGV[1] = 过期阈值时间戳 (now - ttl) +-- ARGV[2] = 目标 userID(传 "" 则不计算 perUserClients) +-- 返回: {totalMembers, uniqueUsers, uniqueClients, perUserClients} +redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]) + +local members = redis.call('ZRANGE', KEYS[1], 0, -1) +local total = #members +if total == 0 then + return {0, 0, 0, 0} +end + +local target_user = ARGV[2] +local users = {} +local clients = {} +local user_count = 0 +local client_count = 0 +local per_user_count = 0 + +for _, member in ipairs(members) do + local sep = string.find(member, '/', 1, true) + if sep then + local uid = string.sub(member, 1, sep - 1) + local cid = string.sub(member, sep + 1) + if not users[uid] then + users[uid] = true + user_count = user_count + 1 + end + if not clients[cid] then + clients[cid] = true + client_count = client_count + 1 + end + if target_user ~= '' and uid == target_user then + per_user_count = per_user_count + 1 + end + end +end + +return {total, user_count, client_count, per_user_count} diff --git a/backend/internal/repository/lua/update_affinity.lua b/backend/internal/repository/lua/update_affinity.lua index 27fe8cb217..6e9b38fcb8 100644 --- a/backend/internal/repository/lua/update_affinity.lua +++ b/backend/internal/repository/lua/update_affinity.lua @@ -1,11 +1,11 @@ -- 原子双写正向+反向索引 --- KEYS[1] = client_affinity:{groupID}:{clientID} (正向: client → accounts) --- KEYS[2] = client_affinity_rev:{groupID}:{accountID} (反向: account → clients) +-- KEYS[1] = affinity:{groupID}:{userID}/{clientID} (正向: member → accounts) +-- KEYS[2] = affinity_rev:{groupID}:{accountID} (反向: account → members) -- ARGV[1] = 当前时间戳 (score) -- ARGV[2] = TTL 秒数 -- ARGV[3] = accountID (正向索引的成员) -- ARGV[4] = 过期阈值时间戳 (now - ttl) --- ARGV[5] = clientID (反向索引的成员) +-- ARGV[5] = {userID}/{clientID} (反向索引的成员) redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[4]) redis.call('ZADD', KEYS[1], ARGV[1], ARGV[3]) redis.call('EXPIRE', KEYS[1], ARGV[2]) diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index 7f17637f8f..ab8e98ad07 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -263,6 +263,7 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers) { accounts.POST("/:id/clear-rate-limit", h.Admin.Account.ClearRateLimit) accounts.POST("/:id/reset-quota", h.Admin.Account.ResetQuota) accounts.GET("/:id/affinity-clients", h.Admin.Account.GetAffinityClients) + accounts.GET("/:id/affinity-details", h.Admin.Account.GetAffinityDetails) accounts.GET("/:id/temp-unschedulable", h.Admin.Account.GetTempUnschedulable) accounts.DELETE("/:id/temp-unschedulable", h.Admin.Account.ClearTempUnschedulable) accounts.POST("/:id/schedulable", h.Admin.Account.SetSchedulable) diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index e25446a16d..976d940e21 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -1204,16 +1204,21 @@ func (a *Account) IsSessionIDMaskingEnabled() bool { return false } -// IsClientAffinityEnabled 检查是否启用客户端亲和调度 -// 仅适用于 Anthropic 账号(OAuth/SetupToken/APIKey) -// 启用后,新会话会优先调度到之前使用过的账号 -func (a *Account) IsClientAffinityEnabled() bool { +// IsAffinityEnabled 检查是否启用亲和调度(统一入口) +// 仅适用于 Anthropic 账号,同时检查新字段 affinity_enabled 和旧字段 client_affinity_enabled(向后兼容) +func (a *Account) IsAffinityEnabled() bool { if a.Platform != PlatformAnthropic { return false } if a.Extra == nil { return false } + if v, ok := a.Extra["affinity_enabled"]; ok { + if enabled, ok := v.(bool); ok { + return enabled + } + return false + } if v, ok := a.Extra["client_affinity_enabled"]; ok { if enabled, ok := v.(bool); ok { return enabled @@ -1222,6 +1227,11 @@ func (a *Account) IsClientAffinityEnabled() bool { return false } +// IsClientAffinityEnabled 向后兼容别名,内部调用 IsAffinityEnabled +func (a *Account) IsClientAffinityEnabled() bool { + return a.IsAffinityEnabled() +} + // AffinityZone 表示账号的客户端亲和分区 type AffinityZone int @@ -1288,6 +1298,150 @@ func (a *Account) GetAffinityZone(clientCount int64) AffinityZone { return AffinityZoneRed } +// IsAffinityAllowSwitch 检查亲和账号是否允许在无绿区时切换到其他账号 +// 默认 true(允许切换),设为 false 时即使全部红区也不切换 +func (a *Account) IsAffinityAllowSwitch() bool { + if a.Extra == nil { + return true + } + if v, ok := a.Extra["affinity_allow_switch"]; ok { + if allow, ok := v.(bool); ok { + return allow + } + } + return true +} + +// GetPinnedUsers 获取指定亲和用户列表 +func (a *Account) GetPinnedUsers() []int64 { + if a.Extra == nil { + return nil + } + v, ok := a.Extra["pinned_users"] + if !ok || v == nil { + return nil + } + arr, ok := v.([]any) + if !ok || len(arr) == 0 { + return nil + } + result := make([]int64, 0, len(arr)) + for _, item := range arr { + switch id := item.(type) { + case float64: + result = append(result, int64(id)) + case int64: + result = append(result, id) + case int: + result = append(result, int64(id)) + case json.Number: + if i, err := id.Int64(); err == nil { + result = append(result, i) + } + } + } + return result +} + +// IsPinnedUser 检查 userID 是否在指定亲和用户列表中 +func (a *Account) IsPinnedUser(userID int64) bool { + for _, id := range a.GetPinnedUsers() { + if id == userID { + return true + } + } + return false +} + +// GetAffinityUserBase 获取用户维度亲和基础限制(绿区上限),0 表示未配置 +func (a *Account) GetAffinityUserBase() int { + if a.Extra == nil { + return 0 + } + if v, ok := a.Extra["affinity_user_base"]; ok { + return parseExtraInt(v) + } + return 0 +} + +// GetAffinityUserBuffer 获取用户维度亲和缓冲区大小(黄区范围) +// 返回 (value, configured),语义与 GetAffinityBuffer 相同 +func (a *Account) GetAffinityUserBuffer() (int, bool) { + if a.Extra == nil { + return 0, false + } + v, ok := a.Extra["affinity_user_buffer"] + if !ok { + return 0, false + } + if v == nil { + return 0, false + } + return parseExtraInt(v), true +} + +// GetPerUserClientLimit 获取每用户客户端限制,0 表示不限制 +func (a *Account) GetPerUserClientLimit() int { + if a.Extra == nil { + return 0 + } + if v, ok := a.Extra["per_user_client_limit"]; ok { + return parseExtraInt(v) + } + return 0 +} + +// getAffinityZoneForDim 根据单一维度的计数和三区参数计算亲和分区 +func getAffinityZoneForDim(count int64, base int, buffer int, bufferConfigured bool) AffinityZone { + if base <= 0 { + return AffinityZoneGreen + } + if count <= int64(base) { + return AffinityZoneGreen + } + if !bufferConfigured { + return AffinityZoneYellow + } + if buffer == 0 { + return AffinityZoneRed + } + if count <= int64(base+buffer) { + return AffinityZoneYellow + } + return AffinityZoneRed +} + +// GetMultiDimAffinityZone 根据多维度计数计算亲和分区,取所有维度中最严格的区域。 +// 未开启亲和的账号永远返回绿区。 +func (a *Account) GetMultiDimAffinityZone(userCount, clientCount, perUserCount int64) AffinityZone { + if !a.IsAffinityEnabled() { + return AffinityZoneGreen + } + worst := AffinityZoneGreen + + // 客户端维度 + clientBase := a.GetAffinityBase() + clientBuffer, clientBufCfg := a.GetAffinityBuffer() + if z := getAffinityZoneForDim(clientCount, clientBase, clientBuffer, clientBufCfg); z > worst { + worst = z + } + + // 用户维度 + userBase := a.GetAffinityUserBase() + userBuffer, userBufCfg := a.GetAffinityUserBuffer() + if z := getAffinityZoneForDim(userCount, userBase, userBuffer, userBufCfg); z > worst { + worst = z + } + + // 每用户客户端维度 + perUserLimit := a.GetPerUserClientLimit() + if perUserLimit > 0 && perUserCount > int64(perUserLimit) { + worst = AffinityZoneRed + } + + return worst +} + // IsCacheTTLOverrideEnabled 检查是否启用缓存 TTL 强制替换 // 仅适用于 Anthropic OAuth/SetupToken 类型账号 // 启用后将所有 cache creation tokens 归入指定的 TTL 类型(5m 或 1h) diff --git a/backend/internal/service/antigravity_smart_retry_test.go b/backend/internal/service/antigravity_smart_retry_test.go index f0e3016bd3..2c20559094 100644 --- a/backend/internal/service/antigravity_smart_retry_test.go +++ b/backend/internal/service/antigravity_smart_retry_test.go @@ -30,12 +30,15 @@ func (c *stubSmartRetryCache) DeleteSessionAccountID(_ context.Context, groupID c.deleteCalls = append(c.deleteCalls, deleteSessionCall{groupID: groupID, sessionHash: sessionHash}) return nil } -func (c *stubSmartRetryCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) { +func (c *stubSmartRetryCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { return nil, nil } -func (c *stubSmartRetryCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { +func (c *stubSmartRetryCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { return nil } +func (c *stubSmartRetryCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} func (c *stubSmartRetryCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { return map[int64]int64{}, nil } diff --git a/backend/internal/service/gateway_affinity_flow.go b/backend/internal/service/gateway_affinity_flow.go new file mode 100644 index 0000000000..cf6e7e6e26 --- /dev/null +++ b/backend/internal/service/gateway_affinity_flow.go @@ -0,0 +1,239 @@ +package service + +import "context" + +// gatewayAffinityFlow encapsulates affinity-specific scheduling steps so the +// main account selection flow can stay focused on generic scheduling. +type gatewayAffinityFlow struct { + svc *GatewayService + ctx context.Context + groupID *int64 + sessionHash string + requestedModel string + affinityClientID string + affinityUserID int64 + platform string + useMixed bool + accountByID map[int64]*Account + isExcluded func(int64) bool +} + +type affinityWaitCandidate struct { + account *Account +} + +func newGatewayAffinityFlow( + svc *GatewayService, + ctx context.Context, + groupID *int64, + sessionHash string, + requestedModel string, + affinityClientID string, + affinityUserID int64, + platform string, + useMixed bool, + accountByID map[int64]*Account, + isExcluded func(int64) bool, +) *gatewayAffinityFlow { + return &gatewayAffinityFlow{ + svc: svc, + ctx: ctx, + groupID: groupID, + sessionHash: sessionHash, + requestedModel: requestedModel, + affinityClientID: affinityClientID, + affinityUserID: affinityUserID, + platform: platform, + useMixed: useMixed, + accountByID: accountByID, + isExcluded: isExcluded, + } +} + +// shouldFilterAccountWithoutClientID excludes affinity-enabled Anthropic accounts +// when metadata.user_id does not provide a usable client_id. +func shouldFilterAccountWithoutClientID(account *Account, affinityClientID string) bool { + if account == nil || affinityClientID != "" { + return false + } + if account.Platform != PlatformAnthropic { + return false + } + return account.IsAffinityEnabled() +} + +func filterAccountsWithoutClientID(accounts []Account, affinityClientID string) []Account { + if affinityClientID != "" { + return accounts + } + filtered := make([]Account, 0, len(accounts)) + for _, acc := range accounts { + if shouldFilterAccountWithoutClientID(&acc, affinityClientID) { + continue + } + filtered = append(filtered, acc) + } + return filtered +} + +func (f *gatewayAffinityFlow) preprocessPinnedUsers(accounts []Account) { + if f.affinityUserID <= 0 || f.affinityClientID == "" || f.svc.cache == nil { + return + } + for i := range accounts { + if accounts[i].IsPinnedUser(f.affinityUserID) && accounts[i].IsAffinityEnabled() { + _ = f.svc.cache.UpdateAffinity( + f.ctx, + derefGroupID(f.groupID), + f.affinityUserID, + f.affinityClientID, + accounts[i].ID, + ClientAffinityTTL, + ) + } + } +} + +// trySelectAffinityAccount runs Layer 1.4 and returns: +// - result != nil: affinity path selected an account or wait plan +// - affinityHit == true: an effective affinity-enabled record was considered and should suppress sticky fallback +func (f *gatewayAffinityFlow) trySelectAffinityAccount() (*AccountSelectionResult, bool, error) { + if f.affinityClientID == "" || f.affinityUserID <= 0 || f.svc.cache == nil { + return nil, false, nil + } + + gid := derefGroupID(f.groupID) + affinityAccountIDs, err := f.svc.cache.GetAffinityAccounts(f.ctx, gid, f.affinityUserID, f.affinityClientID, ClientAffinityTTL) + if err != nil || len(affinityAccountIDs) == 0 { + return nil, false, nil + } + + noSwitchBlocked := false + anyAllowSwitch := false + effectiveAffinityHit := false + waitCandidates := make(map[int64]*affinityWaitCandidate) + + for _, affinityAccID := range affinityAccountIDs { + account, ok := f.accountByID[affinityAccID] + + if f.isExcluded != nil && f.isExcluded(affinityAccID) { + checkAcc := account + if !ok && f.svc.accountRepo != nil { + if acc, repoErr := f.svc.accountRepo.GetByID(f.ctx, affinityAccID); repoErr == nil && acc != nil { + checkAcc = acc + } + } + if checkAcc != nil && checkAcc.IsAffinityEnabled() { + effectiveAffinityHit = true + if !checkAcc.IsAffinityAllowSwitch() { + noSwitchBlocked = true + } else { + anyAllowSwitch = true + } + } + continue + } + + if !ok || !f.svc.isAccountSchedulableForSelection(account) { + checkAcc := account + if !ok && f.svc.accountRepo != nil { + if acc, repoErr := f.svc.accountRepo.GetByID(f.ctx, affinityAccID); repoErr == nil && acc != nil { + checkAcc = acc + } + } + if checkAcc != nil && checkAcc.IsAffinityEnabled() { + effectiveAffinityHit = true + if !checkAcc.IsAffinityAllowSwitch() { + noSwitchBlocked = true + } else { + anyAllowSwitch = true + } + } + continue + } + + if !account.IsAffinityEnabled() { + continue + } + effectiveAffinityHit = true + if account.IsAffinityAllowSwitch() { + anyAllowSwitch = true + } + + if !f.svc.isAccountAllowedForPlatform(account, f.platform, f.useMixed) || + (f.requestedModel != "" && !f.svc.isModelSupportedByAccountWithContext(f.ctx, account, f.requestedModel)) || + !f.svc.isAccountSchedulableForModelSelection(f.ctx, account, f.requestedModel) || + !f.svc.isAccountSchedulableForQuota(account) || + !f.svc.isAccountSchedulableForWindowCost(f.ctx, account, false) || + !f.svc.isAccountSchedulableForRPM(f.ctx, account, false) { + if !account.IsAffinityAllowSwitch() { + noSwitchBlocked = true + } + continue + } + + userCount, clientCount, perUserCount, multiErr := f.svc.cache.GetAffinityMultiCount( + f.ctx, gid, affinityAccID, f.affinityUserID, ClientAffinityTTL, + ) + if multiErr == nil { + zone := account.GetMultiDimAffinityZone(userCount, clientCount, perUserCount) + if zone == AffinityZoneRed { + if !account.IsAffinityAllowSwitch() { + noSwitchBlocked = true + } + continue + } + } + + result, acquireErr := f.svc.tryAcquireAccountSlot(f.ctx, affinityAccID, account.Concurrency) + if acquireErr == nil && result.Acquired { + if !f.svc.checkAndRegisterSession(f.ctx, account, f.sessionHash) { + result.ReleaseFunc() + continue + } + _ = f.svc.cache.UpdateAffinity(f.ctx, gid, f.affinityUserID, f.affinityClientID, affinityAccID, ClientAffinityTTL) + if f.sessionHash != "" { + _ = f.svc.cache.SetSessionAccountID(f.ctx, gid, f.sessionHash, affinityAccID, stickySessionTTL) + } + return &AccountSelectionResult{ + Account: account, + Acquired: true, + ReleaseFunc: result.ReleaseFunc, + }, true, nil + } + if acquireErr == nil && !result.Acquired && !account.IsAffinityAllowSwitch() { + noSwitchBlocked = true + waitCandidates[affinityAccID] = &affinityWaitCandidate{account: account} + } + } + + if noSwitchBlocked && !anyAllowSwitch && f.svc.concurrencyService != nil { + for _, waitAccID := range affinityAccountIDs { + candidate, ok := waitCandidates[waitAccID] + if !ok || candidate == nil || candidate.account == nil { + continue + } + acc := candidate.account + waitingCount, _ := f.svc.concurrencyService.GetAccountWaitingCount(f.ctx, waitAccID) + if waitingCount >= f.svc.schedulingConfig().StickySessionMaxWaiting { + continue + } + if !f.svc.checkAndRegisterSession(f.ctx, acc, f.sessionHash) { + continue + } + cfg := f.svc.schedulingConfig() + return &AccountSelectionResult{ + Account: acc, + WaitPlan: &AccountWaitPlan{ + AccountID: waitAccID, + MaxConcurrency: acc.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }, + }, true, nil + } + return nil, effectiveAffinityHit, ErrAffinityNoSwitch + } + + return nil, effectiveAffinityHit, nil +} diff --git a/backend/internal/service/gateway_affinity_scheduling_test.go b/backend/internal/service/gateway_affinity_scheduling_test.go index 0650081195..f5bf974c0a 100644 --- a/backend/internal/service/gateway_affinity_scheduling_test.go +++ b/backend/internal/service/gateway_affinity_scheduling_test.go @@ -9,6 +9,7 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -36,12 +37,15 @@ func (m *mockAffinityCache) RefreshSessionTTL(_ context.Context, _ int64, _ stri func (m *mockAffinityCache) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { return nil } -func (m *mockAffinityCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) { +func (m *mockAffinityCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { return nil, nil } -func (m *mockAffinityCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { +func (m *mockAffinityCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { return nil } +func (m *mockAffinityCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} func (m *mockAffinityCache) GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) { m.getCountBatchCalls++ if m.getCountBatchFunc != nil { @@ -613,6 +617,28 @@ func TestAffinityIsClientAffinityEnabled(t *testing.T) { } assert.False(t, acc.IsClientAffinityEnabled(), "string 'true' should not enable affinity") }) + + t.Run("new flag false overrides legacy flag true", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "affinity_enabled": false, + "client_affinity_enabled": true, + }, + } + assert.False(t, acc.IsClientAffinityEnabled()) + }) + + t.Run("new flag true overrides legacy flag false", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "affinity_enabled": true, + "client_affinity_enabled": false, + }, + } + assert.True(t, acc.IsClientAffinityEnabled()) + }) } // =========================================================================== @@ -798,3 +824,936 @@ func TestClassifyByAffinityZone(t *testing.T) { require.Len(t, result, 2) }) } + +// =========================================================================== +// GetMultiDimAffinityZone 测试 +// =========================================================================== + +func TestGetMultiDimAffinityZone(t *testing.T) { + makeAccount := func(opts map[string]any) *Account { + extra := map[string]any{"affinity_enabled": true} + for k, v := range opts { + extra[k] = v + } + return &Account{ + Platform: PlatformAnthropic, + Extra: extra, + } + } + + t.Run("user green + client green = green", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "affinity_user_base": 3, + "affinity_user_buffer": 2, + }) + assert.Equal(t, AffinityZoneGreen, acc.GetMultiDimAffinityZone(2, 5, 0)) + }) + + t.Run("user green + client yellow = yellow", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "affinity_user_base": 3, + "affinity_user_buffer": 2, + }) + assert.Equal(t, AffinityZoneYellow, acc.GetMultiDimAffinityZone(2, 12, 0)) + }) + + t.Run("user yellow + client green = yellow", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "affinity_user_base": 3, + "affinity_user_buffer": 2, + }) + assert.Equal(t, AffinityZoneYellow, acc.GetMultiDimAffinityZone(4, 5, 0)) + }) + + t.Run("user red + client green = red", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "affinity_user_base": 3, + "affinity_user_buffer": 0, // no yellow, direct red + }) + assert.Equal(t, AffinityZoneRed, acc.GetMultiDimAffinityZone(4, 5, 0)) + }) + + t.Run("user green + client red = red", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 0, // no yellow, direct red + "affinity_user_base": 3, + "affinity_user_buffer": 2, + }) + assert.Equal(t, AffinityZoneRed, acc.GetMultiDimAffinityZone(2, 11, 0)) + }) + + t.Run("user_base=0 only checks client dimension", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + // no affinity_user_base → 0 → user dimension always green + }) + // client in yellow range + assert.Equal(t, AffinityZoneYellow, acc.GetMultiDimAffinityZone(100, 12, 0)) + }) + + t.Run("perUserCount exceeds limit = red", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "per_user_client_limit": 3, + }) + // client green, user green, but perUser > limit + assert.Equal(t, AffinityZoneRed, acc.GetMultiDimAffinityZone(1, 5, 4)) + }) + + t.Run("perUserCount within limit = green", func(t *testing.T) { + acc := makeAccount(map[string]any{ + "affinity_base": 10, + "affinity_buffer": 5, + "per_user_client_limit": 3, + }) + assert.Equal(t, AffinityZoneGreen, acc.GetMultiDimAffinityZone(1, 5, 3)) + }) + + t.Run("affinity disabled always green", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "affinity_enabled": false, + "affinity_base": 1, + "affinity_buffer": 0, + }, + } + assert.Equal(t, AffinityZoneGreen, acc.GetMultiDimAffinityZone(100, 100, 100)) + }) +} + +// =========================================================================== +// IsAffinityAllowSwitch 测试 +// =========================================================================== + +func TestIsAffinityAllowSwitch(t *testing.T) { + t.Run("default true when no field in Extra", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"affinity_enabled": true}, + } + assert.True(t, acc.IsAffinityAllowSwitch()) + }) + + t.Run("explicit false", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"affinity_allow_switch": false}, + } + assert.False(t, acc.IsAffinityAllowSwitch()) + }) + + t.Run("Extra nil defaults to true", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: nil, + } + assert.True(t, acc.IsAffinityAllowSwitch()) + }) + + t.Run("explicit true", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{"affinity_allow_switch": true}, + } + assert.True(t, acc.IsAffinityAllowSwitch()) + }) +} + +// =========================================================================== +// GetPinnedUsers / IsPinnedUser 测试 +// =========================================================================== + +func TestGetPinnedUsers(t *testing.T) { + t.Run("normal list", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "pinned_users": []any{float64(10), float64(20), float64(30)}, + }, + } + result := acc.GetPinnedUsers() + assert.Equal(t, []int64{10, 20, 30}, result) + }) + + t.Run("empty list", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "pinned_users": []any{}, + }, + } + result := acc.GetPinnedUsers() + assert.Nil(t, result) + }) + + t.Run("Extra nil", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: nil, + } + result := acc.GetPinnedUsers() + assert.Nil(t, result) + }) + + t.Run("pinned_users not set", func(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{}, + } + result := acc.GetPinnedUsers() + assert.Nil(t, result) + }) +} + +func TestIsPinnedUser(t *testing.T) { + acc := &Account{ + Platform: PlatformAnthropic, + Extra: map[string]any{ + "pinned_users": []any{float64(10), float64(20), float64(30)}, + }, + } + + t.Run("hit", func(t *testing.T) { + assert.True(t, acc.IsPinnedUser(20)) + }) + + t.Run("miss", func(t *testing.T) { + assert.False(t, acc.IsPinnedUser(99)) + }) + + t.Run("nil Extra", func(t *testing.T) { + nilAcc := &Account{Platform: PlatformAnthropic, Extra: nil} + assert.False(t, nilAcc.IsPinnedUser(10)) + }) +} + +// =========================================================================== +// Enhanced mock for integration tests (supports affinity lookups + tracking) +// =========================================================================== + +type affinityIntegrationCache struct { + mockAffinityCache + getAffinityAccountsFunc func(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) + updateAffinityCalls []affinityUpdateCall + getMultiCountFunc func(ctx context.Context, groupID int64, accountID int64, targetUserID int64, ttl time.Duration) (int64, int64, int64, error) + sessionBindings map[string]int64 + getAffinityAccountsCalled bool // 标记 GetAffinityAccounts 是否已被调用,用于区分预处理/Layer 2 阶段 +} + +type affinityUpdateCall struct { + groupID int64 + userID int64 + clientID string + accountID int64 + beforeAffinityLookup bool // true 表示该调用发生在 GetAffinityAccounts 之前(预处理阶段) +} + +func (c *affinityIntegrationCache) GetSessionAccountID(_ context.Context, _ int64, hash string) (int64, error) { + if c.sessionBindings != nil { + if id, ok := c.sessionBindings[hash]; ok { + return id, nil + } + } + return 0, errors.New("not found") +} + +func (c *affinityIntegrationCache) GetAffinityAccounts(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) { + c.getAffinityAccountsCalled = true + if c.getAffinityAccountsFunc != nil { + return c.getAffinityAccountsFunc(ctx, groupID, userID, clientID, ttl) + } + return nil, nil +} + +func (c *affinityIntegrationCache) UpdateAffinity(_ context.Context, groupID int64, userID int64, clientID string, accountID int64, _ time.Duration) error { + c.updateAffinityCalls = append(c.updateAffinityCalls, affinityUpdateCall{ + groupID: groupID, userID: userID, clientID: clientID, accountID: accountID, + beforeAffinityLookup: !c.getAffinityAccountsCalled, + }) + return nil +} + +func (c *affinityIntegrationCache) GetAffinityMultiCount(ctx context.Context, groupID int64, accountID int64, targetUserID int64, ttl time.Duration) (int64, int64, int64, error) { + if c.getMultiCountFunc != nil { + return c.getMultiCountFunc(ctx, groupID, accountID, targetUserID, ttl) + } + return 0, 0, 0, nil +} + +// =========================================================================== +// TestAffinityPreprocessPinnedUsers — 验证 pinned_users 预处理逻辑 +// =========================================================================== + +func TestAffinityPreprocessPinnedUsers(t *testing.T) { + ctx := context.Background() + testUserID := int64(42) + testClientID := "user_" + "a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2" + "_account_test" + + t.Run("pinned user triggers UpdateAffinity", func(t *testing.T) { + cache := &affinityIntegrationCache{} + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "pinned_users": []any{float64(42)}, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{"affinity_enabled": true}, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{1: true, 2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + + // 验证 UpdateAffinity 至少被调用了一次(pinned preprocess for account 1) + found := false + for _, call := range cache.updateAffinityCalls { + if call.accountID == 1 && call.userID == testUserID { + found = true + break + } + } + assert.True(t, found, "UpdateAffinity should be called for pinned account 1 with userID %d", testUserID) + }) + + t.Run("non-pinned user does not trigger UpdateAffinity preprocess", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return nil, nil // 无亲和记录,直接进入 Layer 2 + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "pinned_users": []any{float64(99)}, // different user + }, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{1: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + + // 验证预处理阶段(beforeAffinityLookup=true)没有为 userID=42 在 account 1 上调用 UpdateAffinity + for _, call := range cache.updateAffinityCalls { + if call.beforeAffinityLookup && call.accountID == 1 && call.userID == testUserID { + t.Errorf("UpdateAffinity should NOT be called in preprocess for non-pinned user %d on account %d", testUserID, call.accountID) + } + } + }) +} + +// =========================================================================== +// TestAffinityNoSwitchError — 验证 ErrAffinityNoSwitch 行为 +// =========================================================================== + +func TestAffinityNoSwitchError(t *testing.T) { + ctx := context.Background() + testUserID := int64(42) + testClientID := "user_" + "a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2" + "_account_test" + futureTime := time.Now().Add(1 * time.Hour) + + t.Run("affinity hit but unschedulable + allow_switch=false returns ErrAffinityNoSwitch", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, + RateLimitResetAt: &futureTime, // rate limited = unschedulable for selection + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.Error(t, err) + assert.ErrorIs(t, err, ErrAffinityNoSwitch) + }) + + t.Run("affinity hit but unschedulable + allow_switch=true continues to Layer 2", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, + RateLimitResetAt: &futureTime, // rate limited = unschedulable + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": true, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, int64(2), result.Account.ID, "should fall through to Layer 2 and select account 2") + }) + + t.Run("no affinity records - not affected by allow_switch", func(t *testing.T) { + cache := &affinityIntegrationCache{} // returns nil from GetAffinityAccounts + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{1: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, int64(1), result.Account.ID, "should select normally via Layer 2") + }) + + t.Run("affinity disabled account ignores allow_switch=false", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": false, + "affinity_allow_switch": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 2, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + assert.NotEqual(t, int64(0), result.Account.ID) + }) + + t.Run("allow_switch=false + model_unsupported returns ErrAffinityNoSwitch", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + // model_mapping 仅包含 claude-3-haiku → 请求 claude-3-5-sonnet 将被拒绝 + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "claude-3-haiku-20240307": "claude-3-haiku-20240307", + }, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.Error(t, err) + assert.ErrorIs(t, err, ErrAffinityNoSwitch) + }) + + t.Run("allow_switch=false + red_zone returns ErrAffinityNoSwitch", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + getMultiCountFunc: func(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + // userCount=5, clientCount=5, perUserCount=0 + // 账号 affinity_base=1, affinity_buffer=0 → clientCount(5) > base(1) + buffer(0) → 红区 + return 5, 5, 0, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + "affinity_base": 1, + "affinity_buffer": 0, // buffer=0 → 超过 base 直接红区 + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.Error(t, err) + assert.ErrorIs(t, err, ErrAffinityNoSwitch) + }) + + t.Run("one_allow_switch=true overrides another allow_switch=false (一票放行)", func(t *testing.T) { + // 账号 1: allow_switch=false + rate limited(不可调度) + // 账号 3: allow_switch=true + rate limited(不可调度) + // 一票放行:账号 3 允许切换 → 不阻断 → 降级到 Layer 2 选中账号 2 + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1, 3}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, + RateLimitResetAt: &futureTime, // rate limited + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + { + ID: 3, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, + RateLimitResetAt: &futureTime, // rate limited + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": true, + }, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, int64(2), result.Account.ID, "一票放行: should fall through to Layer 2") + }) + + t.Run("excluded affinity account with allow_switch=false returns ErrAffinityNoSwitch", func(t *testing.T) { + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_allow_switch": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + excluded := map[int64]struct{}{1: {}} + _, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", excluded, testClientID, testUserID) + require.Error(t, err) + assert.ErrorIs(t, err, ErrAffinityNoSwitch) + }) + + t.Run("disabled affinity record does not suppress sticky fallback", func(t *testing.T) { + sessionHash := "sticky-disabled-affinity" + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + sessionBindings: map[string]int64{sessionHash: 2}, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": false, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, sessionHash, "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + assert.Equal(t, int64(2), result.Account.ID) + }) + + t.Run("historical affinity record with disabled account still falls back to sticky account", func(t *testing.T) { + sessionHash := "sticky-disabled-status-affinity" + cache := &affinityIntegrationCache{ + getAffinityAccountsFunc: func(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { + return []int64{1}, nil + }, + sessionBindings: map[string]int64{sessionHash: 2}, + } + + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, Platform: PlatformAnthropic, Priority: 1, + Status: StatusDisabled, Schedulable: true, Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": true, + }, + }, + { + ID: 2, Platform: PlatformAnthropic, Priority: 1, + Status: StatusActive, Schedulable: true, Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + concCache := &mockConcurrencyCache{acquireResults: map[int64]bool{2: true}} + concSvc := NewConcurrencyService(concCache) + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: cfg, + concurrencyService: concSvc, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, sessionHash, "claude-3-5-sonnet-20241022", nil, testClientID, testUserID) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + assert.Equal(t, int64(2), result.Account.ID) + }) +} + +// =========================================================================== +// TestAffinityMultiDimClassify — 验证 classifyByAffinityZone 与多维度计数 +// =========================================================================== + +func TestAffinityMultiDimClassify(t *testing.T) { + t.Run("multi-dim zone classification with user and client dimensions", func(t *testing.T) { + makeAWL := func(id int64, extra map[string]any, count int64) accountWithLoad { + return accountWithLoad{ + account: &Account{ID: id, Platform: PlatformAnthropic, Extra: extra}, + loadInfo: &AccountLoadInfo{AccountID: id}, + affinityCount: count, + } + } + + // Account 1: client green (3 <= 5), classified as green by old classifyByAffinityZone + // Account 2: client red (10 > 8), classified as red + accs := []accountWithLoad{ + makeAWL(1, map[string]any{ + "affinity_enabled": true, + "affinity_base": 5, + "affinity_buffer": 3, + }, 3), + makeAWL(2, map[string]any{ + "affinity_enabled": true, + "affinity_base": 5, + "affinity_buffer": 3, + }, 10), + } + + result := classifyByAffinityZone(accs) + require.Len(t, result, 1) + assert.Equal(t, int64(1), result[0].account.ID, "only green account should remain") + }) + + t.Run("all green returns all", func(t *testing.T) { + makeAWL := func(id int64, count int64) accountWithLoad { + return accountWithLoad{ + account: &Account{ + ID: id, + Platform: PlatformAnthropic, + Extra: map[string]any{ + "affinity_enabled": true, + "affinity_base": 10, + "affinity_buffer": 5, + }, + }, + loadInfo: &AccountLoadInfo{AccountID: id}, + affinityCount: count, + } + } + + accs := []accountWithLoad{ + makeAWL(1, 3), + makeAWL(2, 5), + makeAWL(3, 10), // exactly at base + } + + result := classifyByAffinityZone(accs) + require.Len(t, result, 3) + }) +} diff --git a/backend/internal/service/gateway_hotpath_optimization_test.go b/backend/internal/service/gateway_hotpath_optimization_test.go index 0203c3311f..af108c6fc2 100644 --- a/backend/internal/service/gateway_hotpath_optimization_test.go +++ b/backend/internal/service/gateway_hotpath_optimization_test.go @@ -143,12 +143,15 @@ func (s *stickyGatewayCacheHotpathStub) RefreshSessionTTL(ctx context.Context, g func (s *stickyGatewayCacheHotpathStub) DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error { return nil } -func (s *stickyGatewayCacheHotpathStub) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) { +func (s *stickyGatewayCacheHotpathStub) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { return nil, nil } -func (s *stickyGatewayCacheHotpathStub) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { +func (s *stickyGatewayCacheHotpathStub) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { return nil } +func (s *stickyGatewayCacheHotpathStub) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { return map[int64]int64{}, nil } @@ -750,7 +753,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) { modelsListCacheTTL: time.Minute, } - result, err := svc.SelectAccountWithLoadAwareness(baseCtx, nil, "sess-hash", "", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(baseCtx, nil, "sess-hash", "", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -772,7 +775,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) { ctx := context.WithValue(baseCtx, ctxkey.PrefetchedStickyAccountID, account.ID) ctx = context.WithValue(ctx, ctxkey.PrefetchedStickyGroupID, int64(0)) - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -794,7 +797,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) { ctx := context.WithValue(baseCtx, ctxkey.PrefetchedStickyAccountID, int64(999)) ctx = context.WithValue(ctx, ctxkey.PrefetchedStickyGroupID, int64(77)) - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go index fcec495738..8ec4d24c9a 100644 --- a/backend/internal/service/gateway_multiplatform_test.go +++ b/backend/internal/service/gateway_multiplatform_test.go @@ -235,12 +235,15 @@ func (m *mockGatewayCacheForPlatform) DeleteSessionAccountID(ctx context.Context return nil } -func (m *mockGatewayCacheForPlatform) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) { +func (m *mockGatewayCacheForPlatform) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { return nil, nil } -func (m *mockGatewayCacheForPlatform) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { +func (m *mockGatewayCacheForPlatform) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { return nil } +func (m *mockGatewayCacheForPlatform) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} func (m *mockGatewayCacheForPlatform) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { return map[int64]int64{}, nil } @@ -2050,7 +2053,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, // No concurrency service } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2103,7 +2106,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, // legacy path } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-b", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-b", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2135,7 +2138,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2167,13 +2170,314 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { } excludedIDs := map[int64]struct{}{1: {}} - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", excludedIDs, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", excludedIDs, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) require.Equal(t, int64(2), result.Account.ID, "不应选择被排除的账号") }) + t.Run("无客户端ID时过滤Anthropic OAuth亲和账号-load-aware路径", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(2), result.Account.ID, "无客户端ID时应过滤亲和OAuth账号") + }) + + t.Run("无客户端ID且亲和关闭时-load-aware路径恢复旧逻辑", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "affinity_enabled": false, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(1), result.Account.ID, "关闭亲和后不应再被无clientID过滤") + }) + + t.Run("有客户端ID时不过滤Anthropic OAuth亲和账号", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } + + metadataUserID := "user_0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef_account_test_session_test" + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, metadataUserID, 123) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(1), result.Account.ID, "有客户端ID时应允许亲和OAuth账号参与调度") + }) + + t.Run("过滤对Anthropic全类型生效-SetupToken也受影响", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeSetupToken, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 3, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 3, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = true + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}), + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(3), result.Account.ID, "无客户端ID时,Anthropic SetupToken/OAuth(开启亲和)都应被过滤") + }) + + t.Run("无客户端ID时过滤Anthropic OAuth亲和账号-legacy路径", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": true, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = false + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: nil, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(2), result.Account.ID, "legacy路径也应过滤亲和OAuth账号") + }) + + t.Run("无客户端ID且亲和关闭时-legacy路径恢复旧逻辑", func(t *testing.T) { + repo := &mockAccountRepoForPlatform{ + accounts: []Account{ + { + ID: 1, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 1, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + Extra: map[string]any{ + "client_affinity_enabled": false, + }, + }, + { + ID: 2, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Priority: 2, + Status: StatusActive, + Schedulable: true, + Concurrency: 5, + }, + }, + accountsByID: map[int64]*Account{}, + } + for i := range repo.accounts { + repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i] + } + + cfg := testConfig() + cfg.Gateway.Scheduling.LoadBatchEnabled = false + + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: nil, + } + + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, result.Account) + require.Equal(t, int64(1), result.Account.ID, "关闭亲和后 legacy 路径不应再被无clientID过滤") + }) + t.Run("粘性命中-不调用GetByID", func(t *testing.T) { repo := &mockAccountRepoForPlatform{ accounts: []Account{ @@ -2201,7 +2505,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2237,7 +2541,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2278,7 +2582,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(testCtx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(testCtx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2306,7 +2610,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.Error(t, err) require.Nil(t, result) require.ErrorIs(t, err, ErrNoAvailableAccounts) @@ -2338,7 +2642,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2371,7 +2675,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2409,7 +2713,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -2445,7 +2749,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "legacy", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "legacy", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2504,7 +2808,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -2558,7 +2862,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2612,7 +2916,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2670,7 +2974,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2728,7 +3032,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route-full", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route-full", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -2786,7 +3090,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "fallback", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "fallback", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2823,7 +3127,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -2875,7 +3179,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "gemini", "gemini-2.5-pro", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "gemini", "gemini-2.5-pro", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -2953,7 +3257,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { } excluded := map[int64]struct{}{1: {}} - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", excluded, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", excluded, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -3007,7 +3311,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "gemini-2.5-pro", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "gemini-2.5-pro", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) @@ -3040,7 +3344,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: nil, } - result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.Error(t, err) require.Nil(t, result) require.ErrorIs(t, err, ErrClaudeCodeOnly) @@ -3078,7 +3382,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "wait", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "wait", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.WaitPlan) @@ -3116,7 +3420,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) { concurrencyService: NewConcurrencyService(concurrencyCache), } - result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "missing-load", "claude-3-5-sonnet-20241022", nil, "") + result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "missing-load", "claude-3-5-sonnet-20241022", nil, "", int64(0)) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, result.Account) diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 1204e71314..788a984e26 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -364,6 +364,12 @@ var ErrNoAvailableAccounts = errors.New("no available accounts") // ErrClaudeCodeOnly 表示分组仅允许 Claude Code 客户端访问 var ErrClaudeCodeOnly = errors.New("this group only allows Claude Code clients") +// ErrAffinityNoSwitch 表示亲和账号不可用且不允许切换到其他账号 +var ErrAffinityNoSwitch = errors.New("affinity account unavailable and switching is disabled") + +// ErrAffinityLimitExceeded 表示亲和客户端限制已达上限 +var ErrAffinityLimitExceeded = errors.New("affinity client limit exceeded") + // allowedHeaders 白名单headers(参考CRS项目) var allowedHeaders = map[string]bool{ "accept": true, @@ -405,24 +411,29 @@ type GatewayCache interface { // Delete sticky session binding, used to proactively clean up when account becomes unavailable DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error - // GetClientAffinityAccounts 获取客户端亲和账号列表(按最近使用降序),同时清理过期成员 - GetClientAffinityAccounts(ctx context.Context, groupID int64, clientID string, ttl time.Duration) ([]int64, error) - // UpdateClientAffinity 添加/更新客户端亲和关系(更新 score 为当前时间戳,刷新 key TTL) - UpdateClientAffinity(ctx context.Context, groupID int64, clientID string, accountID int64, ttl time.Duration) error - // GetAccountAffinityCountBatch 批量获取账号的亲和客户端数量(惰性清理过期成员) + // GetAffinityAccounts 获取亲和账号列表(按最近使用降序),同时清理过期成员 + GetAffinityAccounts(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) + // UpdateAffinity 添加/更新亲和关系(更新 score 为当前时间戳,刷新 key TTL) + UpdateAffinity(ctx context.Context, groupID int64, userID int64, clientID string, accountID int64, ttl time.Duration) error + // GetAccountAffinityCountBatch 批量获取账号的亲和成员数量(惰性清理过期成员) GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) - // GetAccountAffinityClientsBatch 批量获取每个账号跨所有分组的亲和客户端列表(去重) + // GetAccountAffinityClientsBatch 批量获取每个账号跨所有分组的亲和成员列表(去重) // accountGroups: map[accountID][]groupID + // 返回值成员格式为 {userID}/{clientID} GetAccountAffinityClientsBatch(ctx context.Context, accountGroups map[int64][]int64, ttl time.Duration) (map[int64][]string, error) // GetAccountAffinityClientsWithScores 获取单个账号跨所有分组的亲和客户端列表(含最后活跃时间) GetAccountAffinityClientsWithScores(ctx context.Context, accountID int64, groupIDs []int64, ttl time.Duration) ([]AffinityClient, error) // ClearAccountAffinity 清除指定账号在所有分组的亲和记录(正向+反向索引) - // 用于账号关闭客户端亲和时立即清理旧绑定 + // 用于账号关闭亲和时立即清理旧绑定 ClearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) error + // GetAffinityMultiCount 获取账号的多维度亲和计数 + // 返回: uniqueUsers, uniqueClients, perUserClients + GetAffinityMultiCount(ctx context.Context, groupID int64, accountID int64, targetUserID int64, ttl time.Duration) (users, clients, perUser int64, err error) } -// AffinityClient 亲和客户端信息(含最后活跃时间) +// AffinityClient 亲和客户端信息(含用户 ID 和最后活跃时间) type AffinityClient struct { + UserID int64 `json:"user_id"` ClientID string `json:"client_id"` LastActive time.Time `json:"last_active"` } @@ -1148,8 +1159,10 @@ func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context } // SelectAccountWithLoadAwareness selects account with load-awareness and wait plan. +// 调度流程文档见 docs/ACCOUNT_SCHEDULING_FLOW.md 。 // metadataUserID: 用于客户端亲和调度,从中提取客户端 ID -func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string) (*AccountSelectionResult, error) { +// sub2apiUserID: 系统用户 ID,用于二维亲和调度 +func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string, sub2apiUserID int64) (*AccountSelectionResult, error) { // 调试日志:记录调度入口参数 excludedIDsList := make([]int64, 0, len(excludedIDs)) for id := range excludedIDs { @@ -1181,6 +1194,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro // 提取客户端 ID(用于客户端亲和调度) affinityClientID := extractClientIDFromMetadata(metadataUserID) + affinityUserID := sub2apiUserID if s.debugModelRoutingEnabled() && requestedModel != "" { groupPlatform := "" @@ -1203,6 +1217,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if err != nil { return nil, err } + if shouldFilterAccountWithoutClientID(account, affinityClientID) { + localExcluded[account.ID] = struct{}{} + continue + } result, err := s.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) if err == nil && result.Acquired { @@ -1264,12 +1282,18 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if err != nil { return nil, err } + accounts = filterAccountsWithoutClientID(accounts, affinityClientID) if len(accounts) == 0 { return nil, ErrNoAvailableAccounts } ctx = s.withWindowCostPrefetch(ctx, accounts) ctx = s.withRPMPrefetch(ctx, accounts) + // 提前构建 accountByID(供 Layer 1 和 Layer 1.5 使用) + accountByID := make(map[int64]*Account, len(accounts)) + for i := range accounts { + accountByID[accounts[i].ID] = &accounts[i] + } isExcluded := func(accountID int64) bool { if excludedIDs == nil { return false @@ -1277,12 +1301,19 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro _, excluded := excludedIDs[accountID] return excluded } - - // 提前构建 accountByID(供 Layer 1 和 Layer 1.5 使用) - accountByID := make(map[int64]*Account, len(accounts)) - for i := range accounts { - accountByID[accounts[i].ID] = &accounts[i] - } + affinityFlow := newGatewayAffinityFlow( + s, + ctx, + groupID, + sessionHash, + requestedModel, + affinityClientID, + affinityUserID, + platform, + useMixed, + accountByID, + isExcluded, + ) // 获取模型路由配置(仅 anthropic 平台) var routingAccountIDs []int64 @@ -1485,8 +1516,8 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if sessionHash != "" && s.cache != nil { _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, item.account.ID, stickySessionTTL) } - if affinityClientID != "" && s.cache != nil && item.account.IsClientAffinityEnabled() { - _ = s.cache.UpdateClientAffinity(ctx, derefGroupID(groupID), affinityClientID, item.account.ID, ClientAffinityTTL) + if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && item.account.IsAffinityEnabled() { + _ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, item.account.ID, ClientAffinityTTL) } if s.debugModelRoutingEnabled() { logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) @@ -1525,14 +1556,27 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro } } - // ============ Layer 1.5: 粘性会话(仅在无模型路由配置时生效) ============ - if len(routingAccountIDs) == 0 && sessionHash != "" && stickyAccountID > 0 && !isExcluded(stickyAccountID) { + // ============ Layer 1.3: 用户亲和预处理(pinned_users 自动注入) ============ + affinityFlow.preprocessPinnedUsers(accounts) + + // ============ Layer 1.4: 客户端亲和调度(优先于粘性会话) ============ + affinityHit := false + if affinityResult, hit, err := affinityFlow.trySelectAffinityAccount(); err != nil { + return nil, err + } else { + affinityHit = hit + if affinityResult != nil { + return affinityResult, nil + } + } + + // ============ Layer 1.5: 粘性会话(仅在无模型路由配置 且 亲和未命中时生效) ============ + if !affinityHit && len(routingAccountIDs) == 0 && sessionHash != "" && stickyAccountID > 0 && !isExcluded(stickyAccountID) { accountID := stickyAccountID if accountID > 0 && !isExcluded(accountID) { account, ok := accountByID[accountID] if ok { // 检查账户是否需要清理粘性会话绑定 - // Check if the account needs sticky session cleanup clearSticky := shouldClearStickySession(account, requestedModel) if clearSticky { _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) @@ -1548,7 +1592,6 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency) if err == nil && result.Acquired { // 会话数量限制检查 - // Session count limit check if !s.checkAndRegisterSession(ctx, account, sessionHash) { result.ReleaseFunc() // 释放槽位,继续到 Layer 2 } else { @@ -1563,10 +1606,8 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, accountID) if waitingCount < cfg.StickySessionMaxWaiting { // 会话数量限制检查(等待计划也需要占用会话配额) - // Session count limit check (wait plan also requires session quota) if !s.checkAndRegisterSession(ctx, account, sessionHash) { // 会话限制已满,继续到 Layer 2 - // Session limit full, continue to Layer 2 } else { return &AccountSelectionResult{ Account: account, @@ -1584,76 +1625,6 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro } } - // ============ Layer 1.6: 客户端亲和(仅在粘性会话未命中时生效) ============ - if affinityClientID != "" && s.cache != nil && stickyAccountID <= 0 { - affinityAccountIDs, err := s.cache.GetClientAffinityAccounts(ctx, derefGroupID(groupID), affinityClientID, ClientAffinityTTL) - if err == nil && len(affinityAccountIDs) > 0 { - for _, affinityAccID := range affinityAccountIDs { - if isExcluded(affinityAccID) { - continue - } - account, ok := accountByID[affinityAccID] - if !ok || !s.isAccountSchedulableForSelection(account) { - continue - } - if !account.IsClientAffinityEnabled() { - continue - } - if !s.isAccountAllowedForPlatform(account, platform, useMixed) { - continue - } - if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, account, requestedModel) { - continue - } - if !s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) { - continue - } - if !s.isAccountSchedulableForQuota(account) { - continue - } - if !s.isAccountSchedulableForWindowCost(ctx, account, false) { - continue - } - if !s.isAccountSchedulableForRPM(ctx, account, false) { - continue - } - // 亲和三区检查:红区账号不可通过亲和命中调度 - if account.GetAffinityBase() > 0 && s.cache != nil { - countMap, err := s.cache.GetAccountAffinityCountBatch(ctx, derefGroupID(groupID), []int64{affinityAccID}, ClientAffinityTTL) - if err == nil { - zone := account.GetAffinityZone(countMap[affinityAccID]) - if zone == AffinityZoneRed { - continue - } - } - } - - result, err := s.tryAcquireAccountSlot(ctx, affinityAccID, account.Concurrency) - if err == nil && result.Acquired { - if !s.checkAndRegisterSession(ctx, account, sessionHash) { - result.ReleaseFunc() - continue - } - // 亲和命中:更新亲和 score + 绑定粘性会话 - _ = s.cache.UpdateClientAffinity(ctx, derefGroupID(groupID), affinityClientID, affinityAccID, ClientAffinityTTL) - if sessionHash != "" { - _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, affinityAccID, stickySessionTTL) - } - slog.Debug("client_affinity_hit", - "group_id", derefGroupID(groupID), - "client_id", affinityClientID[:8]+"...", - "account_id", affinityAccID) - return &AccountSelectionResult{ - Account: account, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil - } - } - // 所有亲和账号不可用,继续到 Layer 2 - } - } - // ============ Layer 2: 负载感知选择 ============ candidates := make([]*Account, 0, len(accounts)) for i := range accounts { @@ -1706,8 +1677,8 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro loadMap, err := s.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) if err != nil { if result, ok := s.tryAcquireByLegacyOrder(ctx, candidates, groupID, sessionHash, preferOAuth); ok { - if affinityClientID != "" && s.cache != nil && result.Account != nil && result.Account.IsClientAffinityEnabled() { - _ = s.cache.UpdateClientAffinity(ctx, derefGroupID(groupID), affinityClientID, result.Account.ID, ClientAffinityTTL) + if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && result.Account != nil && result.Account.IsAffinityEnabled() { + _ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, result.Account.ID, ClientAffinityTTL) } return result, nil } @@ -1771,9 +1742,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if sessionHash != "" && s.cache != nil { _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.account.ID, stickySessionTTL) } - // 更新客户端亲和关系 - if affinityClientID != "" && s.cache != nil && selected.account.IsClientAffinityEnabled() { - _ = s.cache.UpdateClientAffinity(ctx, derefGroupID(groupID), affinityClientID, selected.account.ID, ClientAffinityTTL) + // 更新亲和关系 + if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && selected.account.IsAffinityEnabled() { + _ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, selected.account.ID, ClientAffinityTTL) } return &AccountSelectionResult{ Account: selected.account, @@ -2541,7 +2512,7 @@ func (s *GatewayService) populateAffinityCounts(ctx context.Context, accounts [] // 快速检查:是否有任何账号开启了亲和 hasAffinity := false for _, acc := range accounts { - if acc.account.IsClientAffinityEnabled() { + if acc.account.IsAffinityEnabled() { hasAffinity = true break } @@ -2632,7 +2603,7 @@ func classifyByAffinityZone(accounts []accountWithLoad) []accountWithLoad { // 快速检查:是否有任何账号配置了 affinity_base hasZoneConfig := false for _, acc := range accounts { - if acc.account.IsClientAffinityEnabled() && acc.account.GetAffinityBase() > 0 { + if acc.account.IsAffinityEnabled() && acc.account.GetAffinityBase() > 0 { hasZoneConfig = true break } diff --git a/backend/internal/service/gemini_multiplatform_test.go b/backend/internal/service/gemini_multiplatform_test.go index 7299261429..6972689371 100644 --- a/backend/internal/service/gemini_multiplatform_test.go +++ b/backend/internal/service/gemini_multiplatform_test.go @@ -288,12 +288,15 @@ func (m *mockGatewayCacheForGemini) DeleteSessionAccountID(ctx context.Context, return nil } -func (m *mockGatewayCacheForGemini) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) { +func (m *mockGatewayCacheForGemini) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { return nil, nil } -func (m *mockGatewayCacheForGemini) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { +func (m *mockGatewayCacheForGemini) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { return nil } +func (m *mockGatewayCacheForGemini) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} func (m *mockGatewayCacheForGemini) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { return map[int64]int64{}, nil } diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 889a4cd3e9..2828802fe6 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -282,12 +282,15 @@ func (c *stubGatewayCache) DeleteSessionAccountID(ctx context.Context, groupID i return nil } -func (c *stubGatewayCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) { +func (c *stubGatewayCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { return nil, nil } -func (c *stubGatewayCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { +func (c *stubGatewayCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { return nil } +func (c *stubGatewayCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} func (c *stubGatewayCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { return map[int64]int64{}, nil } diff --git a/backend/internal/service/openai_ws_state_store_test.go b/backend/internal/service/openai_ws_state_store_test.go index 40cce0dd47..9f35cd7ec6 100644 --- a/backend/internal/service/openai_ws_state_store_test.go +++ b/backend/internal/service/openai_ws_state_store_test.go @@ -193,12 +193,15 @@ func (c *openAIWSStateStoreTimeoutProbeCache) DeleteSessionAccountID(ctx context return nil } -func (c *openAIWSStateStoreTimeoutProbeCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) { +func (c *openAIWSStateStoreTimeoutProbeCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { return nil, nil } -func (c *openAIWSStateStoreTimeoutProbeCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { +func (c *openAIWSStateStoreTimeoutProbeCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { return nil } +func (c *openAIWSStateStoreTimeoutProbeCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} func (c *openAIWSStateStoreTimeoutProbeCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { return map[int64]int64{}, nil } diff --git a/backend/internal/service/ops_retry.go b/backend/internal/service/ops_retry.go index fdabbafde9..c0e814ab7b 100644 --- a/backend/internal/service/ops_retry.go +++ b/backend/internal/service/ops_retry.go @@ -519,7 +519,7 @@ func (s *OpsService) selectAccountForRetry(ctx context.Context, reqType opsRetry if s.gatewayService == nil { return nil, fmt.Errorf("gateway service not available") } - return s.gatewayService.SelectAccountWithLoadAwareness(ctx, groupID, "", model, excludedIDs, "") // 重试不使用会话限制 + return s.gatewayService.SelectAccountWithLoadAwareness(ctx, groupID, "", model, excludedIDs, "", int64(0)) // 重试不使用会话限制 default: return nil, fmt.Errorf("unsupported retry type: %s", reqType) } diff --git a/backend/internal/testutil/stubs.go b/backend/internal/testutil/stubs.go index 229985434a..1ae72127bc 100644 --- a/backend/internal/testutil/stubs.go +++ b/backend/internal/testutil/stubs.go @@ -100,12 +100,15 @@ func (c StubGatewayCache) RefreshSessionTTL(_ context.Context, _ int64, _ string func (c StubGatewayCache) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { return nil } -func (c StubGatewayCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) { +func (c StubGatewayCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) { return nil, nil } -func (c StubGatewayCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error { +func (c StubGatewayCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error { return nil } +func (c StubGatewayCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) { + return 0, 0, 0, nil +} func (c StubGatewayCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) { return map[int64]int64{}, nil } diff --git a/docs/ACCOUNT_SCHEDULING_FLOW.md b/docs/ACCOUNT_SCHEDULING_FLOW.md new file mode 100644 index 0000000000..3f7cc807ae --- /dev/null +++ b/docs/ACCOUNT_SCHEDULING_FLOW.md @@ -0,0 +1,61 @@ +# Account Scheduling Flow(SelectAccountWithLoadAwareness) + +本文档对应后端主调度入口 `SelectAccountWithLoadAwareness`,用于说明 Anthropic 网关的账号选择流程、亲和策略与关键过滤规则。 + +## 代码入口与引用 + +- 调度主入口:`backend/internal/service/gateway_service.go` 中的 `SelectAccountWithLoadAwareness` +- 亲和详情 API:`backend/internal/handler/admin/account_affinity_handler.go` 中的 `GetAffinityDetails` +- Admin 路由:`backend/internal/server/routes/admin.go`(`/api/v1/admin/accounts/:id/affinity-details`) + +## 流程概览(修正版) + +1. 解析输入参数 +`sessionHash`、`stickyAccountID`、`affinityClientID(metadata.user_id 提取)`、`affinityUserID(sub2api user id)`。 + +2. Claude Code 限制与分组降级 +先执行 `checkClaudeCodeRestriction`,必要时替换 `groupID`。 + +3. Legacy / Load-aware 分支 +若 `concurrencyService == nil || !LoadBatchEnabled` 走传统路径;否则走负载感知路径。 + +4. Layer 1(模型路由优先) +- 路由候选过滤(排除、平台、模型、配额、窗口费用、RPM) +- 路由范围内 sticky 优先 +- 路由内按(优先级 > 负载 > 亲和数 > LRU)尝试获取槽位 + +5. Layer 1.3(pinned_users 预处理) +仅做 `UpdateAffinity` 预热,不直接做调度决策。 + +6. Layer 1.4(客户端亲和调度) +按亲和记录尝试命中;`allow_switch=false` 且无一票放行时尝试等待计划,否则返回 `ErrAffinityNoSwitch`。 + +7. Layer 1.5(粘性会话) +仅在 `!affinityHit && routingAccountIDs==0` 时生效;先过 `shouldClearStickySession`。 + +8. Layer 2(负载感知选择) +分层过滤:优先级 -> 亲和区(单维客户端) -> 最低负载 -> 最少亲和客户端 -> LRU。 + +9. Layer 3(兜底排队) +按 `FallbackSelectionMode`(`last_used` 或 `random`)排序后返回等待计划。 +注意:此层不是“按最低负载”。 + +## 关键业务规则(2026-03) + +### 无客户端 ID 时过滤亲和 Anthropic 账号 + +当 `metadata.user_id` 无法提取 `client_id`(即 `affinityClientID == ""`)时: + +- 会直接过滤 **Anthropic 平台且开启客户端亲和** 的账号(覆盖 OAuth / SetupToken / API Key / Bedrock 等类型) +- 该规则只看平台与亲和开关,不再限定账号类型 + +设计目标:避免无客户端标识的请求误用客户端亲和账号。 + +## 亲和详情 API 契约 + +`GET /api/v1/admin/accounts/:id/affinity-details` 返回结构(后端已对齐前端): + +- `users[]`:包含 `user_id`、`user_email`、`client_count`、`is_pinned`、`clients[]` +- `total_users` +- `total_clients` +- `pinned_users` diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts index c9584ed5ea..0c3b6f6713 100644 --- a/frontend/src/api/admin/accounts.ts +++ b/frontend/src/api/admin/accounts.ts @@ -17,7 +17,8 @@ import type { AdminDataPayload, AdminDataImportResult, CheckMixedChannelRequest, - CheckMixedChannelResponse + CheckMixedChannelResponse, + AffinityDetailsResponse } from '@/types' /** @@ -630,6 +631,18 @@ export async function getAffinityClients(id: number): Promise<{ client_id: strin return data } +/** + * Get affinity details for an account with user-level grouping + * @param id - Account ID + * @returns Affinity details with user groups + */ +export async function getAffinityDetails(id: number): Promise { + const { data } = await apiClient.get( + `/admin/accounts/${id}/affinity-details` + ) + return data +} + export const accountsAPI = { list, listWithEtag, @@ -667,7 +680,8 @@ export const accountsAPI = { getAntigravityDefaultModelMapping, batchClearError, batchRefresh, - getAffinityClients + getAffinityClients, + getAffinityDetails } export default accountsAPI diff --git a/frontend/src/components/account/AccountCapacityCell.vue b/frontend/src/components/account/AccountCapacityCell.vue index 129e7ef919..aa926d5eab 100644 --- a/frontend/src/components/account/AccountCapacityCell.vue +++ b/frontend/src/components/account/AccountCapacityCell.vue @@ -29,7 +29,7 @@ - + @@ -177,10 +177,11 @@ const rpmTooltip = computed(() => { const showAffinity = computed(() => props.account.platform === 'anthropic' && props.account.client_affinity_enabled === true && - props.account.affinity_client_count != null + (props.account.affinity_client_count != null || props.account.affinity_user_count != null) ) const affinityClientCount = computed(() => props.account.affinity_client_count ?? 0) +const affinityUserCount = computed(() => props.account.affinity_user_count ?? 0) const affinityBase = computed(() => { const extra = props.account.extra as Record | undefined @@ -197,6 +198,21 @@ const affinityBuffer = computed((): number | null => { return null }) +const affinityUserBase = computed(() => { + const extra = props.account.extra as Record | undefined + if (!extra) return 0 + const v = extra.affinity_user_base + return (typeof v === 'number' && v > 0) ? v : 0 +}) + +const affinityUserBuffer = computed((): number | null => { + const extra = props.account.extra as Record | undefined + if (!extra) return null + const v = extra.affinity_user_buffer + if (typeof v === 'number') return v + return null +}) + // 格式化费用显示 const formatCost = (value: number | null | undefined) => { if (value === null || value === undefined) return '0' diff --git a/frontend/src/components/account/AffinityBadge.vue b/frontend/src/components/account/AffinityBadge.vue index ee97fdc4bd..99caf3de76 100644 --- a/frontend/src/components/account/AffinityBadge.vue +++ b/frontend/src/components/account/AffinityBadge.vue @@ -11,9 +11,24 @@ - {{ count }} - / - {{ limitDisplay }} + + + + + + @@ -28,7 +43,7 @@ >
- {{ t('admin.accounts.affinityClients', { count }) }} + {{ t('admin.accounts.affinityDetailTitle') }} @@ -46,26 +61,43 @@
- + +
+ {{ t('admin.accounts.affinityUsers', { count: details.total_users }) }} + {{ t('admin.accounts.affinityClients', { count: details.total_clients }) }} +
+ +
-
+
{{ t('common.loading') }}...
-
- {{ t('admin.accounts.affinityNoClients') }} +
+ {{ t('admin.accounts.affinityNoUsers') }}
-
- - {{ client.client_id }} - - - {{ formatRelativeTime(client.last_active) }} - +
+ +
+
+ P + + {{ userGroup.user_email || `User #${userGroup.user_id}` }} + +
+ + {{ t('admin.accounts.affinityClientCountLabel', { count: userGroup.client_count }) }} + +
+ +
+ + {{ client.client_id }} + + + {{ formatRelativeTime(client.last_active) }} + +
@@ -78,66 +110,76 @@ diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index d92ec70bac..ff6e57e35e 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -1564,7 +1564,7 @@

{{ t('admin.accounts.quotaControl.title') }}

- {{ t('admin.accounts.quotaLimitHint') }} + {{ t('admin.accounts.quotaControl.hint') }}

@@ -1962,9 +1972,19 @@ :enabled="clientAffinityEnabled" :base="affinityBase" :buffer="affinityBuffer" + :allow-switch="affinityAllowSwitch" + :user-base="affinityUserBase" + :user-buffer="affinityUserBuffer" + :per-user-limit="perUserClientLimit" + :pinned-users="pinnedUsers" @update:enabled="clientAffinityEnabled = $event" @update:base="affinityBase = $event" @update:buffer="affinityBuffer = $event" + @update:allow-switch="affinityAllowSwitch = $event" + @update:user-base="affinityUserBase = $event" + @update:user-buffer="affinityUserBuffer = $event" + @update:per-user-limit="perUserClientLimit = $event" + @update:pinned-users="pinnedUsers = $event" /> @@ -3120,6 +3140,11 @@ const showGeminiHelpDialog = ref(false) const clientAffinityEnabled = ref(false) const affinityBase = ref(null) const affinityBuffer = ref(null) +const affinityAllowSwitch = ref(true) +const affinityUserBase = ref(null) +const affinityUserBuffer = ref(null) +const perUserClientLimit = ref(null) +const pinnedUsers = ref([]) // Quota control state (Anthropic OAuth/SetupToken only) const windowCostEnabled = ref(false) @@ -3773,6 +3798,11 @@ const resetForm = () => { clientAffinityEnabled.value = false affinityBase.value = null affinityBuffer.value = null + affinityAllowSwitch.value = true + affinityUserBase.value = null + affinityUserBuffer.value = null + perUserClientLimit.value = null + pinnedUsers.value = [] modelMappings.value = [] modelRestrictionMode.value = 'whitelist' allowedModels.value = [...claudeModels] // Default fill related models @@ -3887,6 +3917,7 @@ const buildAnthropicExtra = (base?: Record): Record) => { if (clientAffinityEnabled.value) { + extra.affinity_enabled = true extra.client_affinity_enabled = true if (affinityBase.value != null && affinityBase.value > 0) { extra.affinity_base = affinityBase.value @@ -3898,10 +3929,38 @@ const applyClientAffinity = (extra: Record) => { } else { delete extra.affinity_buffer } + // v2 fields + extra.affinity_allow_switch = affinityAllowSwitch.value + if (affinityUserBase.value != null && affinityUserBase.value > 0) { + extra.affinity_user_base = affinityUserBase.value + } else { + delete extra.affinity_user_base + } + if (affinityUserBase.value != null && affinityUserBase.value > 0 && affinityUserBuffer.value != null) { + extra.affinity_user_buffer = affinityUserBuffer.value + } else { + delete extra.affinity_user_buffer + } + if (perUserClientLimit.value != null && perUserClientLimit.value > 0) { + extra.per_user_client_limit = perUserClientLimit.value + } else { + delete extra.per_user_client_limit + } + if (pinnedUsers.value.length > 0) { + extra.pinned_users = pinnedUsers.value + } else { + delete extra.pinned_users + } } else { + delete extra.affinity_enabled delete extra.client_affinity_enabled delete extra.affinity_base delete extra.affinity_buffer + delete extra.affinity_allow_switch + delete extra.affinity_user_base + delete extra.affinity_user_buffer + delete extra.per_user_client_limit + delete extra.pinned_users } } diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index 75bbaf39cb..ec4dc9e926 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -1158,7 +1158,7 @@

{{ t('admin.accounts.quotaControl.title') }}

- {{ t('admin.accounts.quotaLimitHint') }} + {{ t('admin.accounts.quotaControl.hint') }}

-
@@ -1902,6 +1921,11 @@ const cacheTTLOverrideTarget = ref('5m') const clientAffinityEnabled = ref(false) const affinityBase = ref(null) const affinityBuffer = ref(null) +const affinityAllowSwitch = ref(true) +const affinityUserBase = ref(null) +const affinityUserBuffer = ref(null) +const perUserClientLimit = ref(null) +const pinnedUsers = ref([]) // OpenAI 自动透传开关(OAuth/API Key) const openaiPassthroughEnabled = ref(false) @@ -2512,6 +2536,11 @@ function loadQuotaControlSettings(account: Account) { clientAffinityEnabled.value = false affinityBase.value = null affinityBuffer.value = null + affinityAllowSwitch.value = true + affinityUserBase.value = null + affinityUserBuffer.value = null + perUserClientLimit.value = null + pinnedUsers.value = [] // Remaining quota control settings only apply to Anthropic accounts if (account.platform !== 'anthropic') { @@ -2530,6 +2559,17 @@ function loadQuotaControlSettings(account: Account) { // buffer: null = infinite yellow, 0 = no yellow, >0 = yellow range const buf = extra.affinity_buffer affinityBuffer.value = (typeof buf === 'number') ? buf : null + + // New v2 fields + affinityAllowSwitch.value = extra.affinity_allow_switch !== false + const ub = extra.affinity_user_base + affinityUserBase.value = (typeof ub === 'number' && ub > 0) ? ub : null + const ubuf = extra.affinity_user_buffer + affinityUserBuffer.value = (typeof ubuf === 'number') ? ubuf : null + const pul = extra.per_user_client_limit + perUserClientLimit.value = (typeof pul === 'number' && pul > 0) ? pul : null + const pu = extra.pinned_users + pinnedUsers.value = Array.isArray(pu) ? pu : [] } // Window cost / session limit only apply to Anthropic OAuth/SetupToken accounts @@ -2951,11 +2991,12 @@ const handleSubmit = async () => { } // For all Anthropic accounts, handle client_affinity in extra - if (props.account.platform === 'anthropic') { - const currentExtra = (props.account.extra as Record) || {} - const newExtra: Record = { ...currentExtra } - if (clientAffinityEnabled.value) { - newExtra.client_affinity_enabled = true + if (props.account.platform === 'anthropic') { + const currentExtra = (props.account.extra as Record) || {} + const newExtra: Record = { ...currentExtra } + if (clientAffinityEnabled.value) { + newExtra.affinity_enabled = true + newExtra.client_affinity_enabled = true if (affinityBase.value != null && affinityBase.value > 0) { newExtra.affinity_base = affinityBase.value } else { @@ -2967,10 +3008,38 @@ const handleSubmit = async () => { } else { delete newExtra.affinity_buffer } - } else { - delete newExtra.client_affinity_enabled - delete newExtra.affinity_base - delete newExtra.affinity_buffer + // v2 fields + newExtra.affinity_allow_switch = affinityAllowSwitch.value + if (affinityUserBase.value != null && affinityUserBase.value > 0) { + newExtra.affinity_user_base = affinityUserBase.value + } else { + delete newExtra.affinity_user_base + } + if (affinityUserBase.value != null && affinityUserBase.value > 0 && affinityUserBuffer.value != null) { + newExtra.affinity_user_buffer = affinityUserBuffer.value + } else { + delete newExtra.affinity_user_buffer + } + if (perUserClientLimit.value != null && perUserClientLimit.value > 0) { + newExtra.per_user_client_limit = perUserClientLimit.value + } else { + delete newExtra.per_user_client_limit + } + if (pinnedUsers.value.length > 0) { + newExtra.pinned_users = pinnedUsers.value + } else { + delete newExtra.pinned_users + } + } else { + newExtra.affinity_enabled = false + newExtra.client_affinity_enabled = false + delete newExtra.affinity_base + delete newExtra.affinity_buffer + delete newExtra.affinity_allow_switch + delete newExtra.affinity_user_base + delete newExtra.affinity_user_buffer + delete newExtra.per_user_client_limit + delete newExtra.pinned_users } updatePayload.extra = newExtra } diff --git a/frontend/src/components/account/__tests__/AffinityBadge.spec.ts b/frontend/src/components/account/__tests__/AffinityBadge.spec.ts index 0dbe186f0b..e7561e7f76 100644 --- a/frontend/src/components/account/__tests__/AffinityBadge.spec.ts +++ b/frontend/src/components/account/__tests__/AffinityBadge.spec.ts @@ -3,7 +3,8 @@ import { mount } from '@vue/test-utils' import AffinityBadge from '../AffinityBadge.vue' vi.mock('@/api/admin/accounts', () => ({ - getAffinityClients: vi.fn() + getAffinityClients: vi.fn(), + getAffinityDetails: vi.fn() })) vi.mock('vue-i18n', async () => { @@ -16,79 +17,147 @@ vi.mock('vue-i18n', async () => { } }) -function mountBadge(count: number, base = 5, buffer: number | null = 10) { +function mountBadge(opts: { + clientCount: number + base?: number + buffer?: number | null + userCount?: number + userBase?: number + userBuffer?: number | null +}) { return mount(AffinityBadge, { props: { accountId: 42, - count, - base, - buffer + clientCount: opts.clientCount, + base: opts.base ?? 5, + buffer: opts.buffer ?? 10, + userCount: opts.userCount ?? 0, + userBase: opts.userBase ?? 0, + userBuffer: opts.userBuffer ?? null } }) } describe('AffinityBadge', () => { - it('renders the correct count number', () => { - const wrapper = mountBadge(5) + // ====== Original tests (client dimension only) ====== + it('renders the correct client count number', () => { + const wrapper = mountBadge({ clientCount: 5 }) expect(wrapper.text()).toContain('5') }) it('renders configured limit text', () => { - const wrapper = mountBadge(5, 5, 10) + const wrapper = mountBadge({ clientCount: 5, base: 5, buffer: 10 }) expect(wrapper.text()).toContain('15') }) it('renders infinity limit text when base is not configured', () => { - const wrapper = mountBadge(5, 0, null) - expect(wrapper.text()).toContain('∞') + const wrapper = mountBadge({ clientCount: 5, base: 0, buffer: null }) + expect(wrapper.text()).toContain('\u221E') }) it('applies red badge class when count exceeds base plus buffer', () => { - const wrapper = mountBadge(16, 5, 10) + const wrapper = mountBadge({ clientCount: 16, base: 5, buffer: 10 }) const badge = wrapper.find('span') expect(badge.classes()).toContain('bg-red-100') expect(badge.classes()).toContain('text-red-700') }) it('applies yellow badge class when count is in buffer range', () => { - const wrapper = mountBadge(6, 5, 10) + const wrapper = mountBadge({ clientCount: 6, base: 5, buffer: 10 }) const badge = wrapper.find('span') expect(badge.classes()).toContain('bg-yellow-100') expect(badge.classes()).toContain('text-yellow-700') }) it('applies yellow badge class when buffer is infinite', () => { - const wrapper = mountBadge(6, 5, null) + const wrapper = mountBadge({ clientCount: 6, base: 5, buffer: null }) const badge = wrapper.find('span') expect(badge.classes()).toContain('bg-yellow-100') expect(badge.classes()).toContain('text-yellow-700') }) it('applies green badge class when count is within base', () => { - const wrapper = mountBadge(5, 5, 10) + const wrapper = mountBadge({ clientCount: 5, base: 5, buffer: 10 }) const badge = wrapper.find('span') expect(badge.classes()).toContain('bg-emerald-100') expect(badge.classes()).toContain('text-emerald-700') }) it('applies gray badge class when count is 0', () => { - const wrapper = mountBadge(0, 5, 10) + const wrapper = mountBadge({ clientCount: 0, base: 5, buffer: 10 }) const badge = wrapper.find('span') expect(badge.classes()).toContain('bg-gray-100') expect(badge.classes()).toContain('text-gray-600') }) it('does NOT show popover initially', () => { - const wrapper = mountBadge(3) + const wrapper = mountBadge({ clientCount: 3 }) expect(wrapper.html()).not.toContain('divide-y') - expect(wrapper.html()).not.toContain('affinityClients') }) it('has mouseenter and mouseleave handlers on the badge', () => { - const wrapper = mountBadge(3) + const wrapper = mountBadge({ clientCount: 3 }) const badge = wrapper.find('span') expect(badge.exists()).toBe(true) badge.trigger('mouseenter') badge.trigger('mouseleave') }) + + // ====== Dual dimension tests ====== + it('dual dimension: user green + client yellow = yellow', () => { + const wrapper = mountBadge({ + clientCount: 6, base: 5, buffer: 10, + userCount: 3, userBase: 5, userBuffer: 5 + }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-yellow-100') + expect(badge.classes()).toContain('text-yellow-700') + }) + + it('dual dimension: user green + client red = red', () => { + const wrapper = mountBadge({ + clientCount: 16, base: 5, buffer: 10, + userCount: 3, userBase: 5, userBuffer: 5 + }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-red-100') + expect(badge.classes()).toContain('text-red-700') + }) + + it('dual dimension: user red + client green = red', () => { + const wrapper = mountBadge({ + clientCount: 3, base: 5, buffer: 10, + userCount: 11, userBase: 5, userBuffer: 5 + }) + const badge = wrapper.find('span') + expect(badge.classes()).toContain('bg-red-100') + expect(badge.classes()).toContain('text-red-700') + }) + + it('user only dimension shows user count text', () => { + const wrapper = mountBadge({ + clientCount: 0, base: 0, buffer: null, + userCount: 3, userBase: 5, userBuffer: 5 + }) + expect(wrapper.text()).toContain('3') + expect(wrapper.text()).toContain('10') + }) + + it('client only dimension shows client count text', () => { + const wrapper = mountBadge({ + clientCount: 4, base: 5, buffer: 10, + userCount: 0, userBase: 0, userBuffer: null + }) + expect(wrapper.text()).toContain('4') + expect(wrapper.text()).toContain('15') + }) + + it('dual dimension limit text shows both U and C prefixes', () => { + const wrapper = mountBadge({ + clientCount: 3, base: 5, buffer: 10, + userCount: 2, userBase: 4, userBuffer: 3 + }) + expect(wrapper.text()).toContain('U2/7') + expect(wrapper.text()).toContain('C3/15') + }) }) diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index 6484a27a28..5f2ef62b4f 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -2177,7 +2177,7 @@ export default { // Quota control (Anthropic OAuth/SetupToken only) quotaControl: { title: 'Quota Control', - hint: 'Configure cost window, session limits, client affinity and other scheduling controls.', + hint: 'Configure affinity settings (user affinity, client affinity), cost windows, session limits and other scheduling controls.', windowCost: { label: '5h Window Cost Limit', hint: 'Limit account cost usage within the 5-hour window', @@ -2238,12 +2238,15 @@ export default { hint: 'When enabled, new sessions prefer accounts previously used by this client to reduce account switching' } }, + affinityConfigTitle: 'Affinity Settings', + affinityConfigHint: 'Configure user affinity and client affinity scheduling rules under quota control.', affinityNoClients: 'No affinity clients', - affinityClients: '{count} affinity clients:', + affinityClients: '{count} client affinity', + affinityClientCountLabel: '{count} clients', affinitySection: 'Client Affinity', affinitySectionHint: 'Control how clients are distributed across accounts. Configure zone thresholds to balance load.', - affinityToggle: 'Enable Client Affinity', - affinityToggleHint: 'New sessions prefer accounts previously used by this client', + affinityToggle: 'Enable Affinity Scheduling', + affinityToggleHint: 'When enabled, new sessions prefer user affinity and client affinity matches', affinityBase: 'Base Limit (Green Zone)', affinityBasePlaceholder: 'Empty = no limit', affinityBaseHint: 'Max clients in green zone (full priority scheduling)', @@ -2252,6 +2255,28 @@ export default { affinityBufferPlaceholder: 'e.g. 3', affinityBufferHint: 'Additional clients allowed in the yellow zone (degraded priority)', affinityBufferInfinite: 'Unlimited', + affinityAllowSwitch: 'Allow Switch', + affinityAllowSwitchHint: 'Allow scheduling to other accounts when affinity account is unavailable', + affinityAllowSwitchWarning: 'When disabled, unavailable affinity account will return error directly', + affinityUserSection: 'User Affinity', + affinityUserSectionHint: 'Configure affinity capacity and scheduling rules for users', + affinityUserBase: 'Base Limit (Green Zone)', + affinityUserBaseHint: 'Max users in green zone (full priority scheduling)', + affinityUserBuffer: 'Buffer (Yellow Zone)', + affinityUserBufferHint: 'Additional users allowed in the yellow zone (degraded priority)', + affinityClientSection: 'Client Affinity', + affinityClientSectionHint: 'Configure affinity capacity and scheduling rules for clients', + affinityPerUserLimit: 'Per-User Client Limit', + affinityPerUserLimitHint: 'Limit the number of distinct clients each user can use', + affinityPerUserMax: 'Max Clients', + affinityPinnedUsers: 'Pinned Affinity Users', + affinityPinnedUsersHint: 'Pre-bound users will automatically create affinity cache and occupy user affinity slots', + affinityPinnedUsersSearch: 'Search user email or username...', + affinityPinnedUsersEmpty: 'No pinned affinity users', + affinityUsers: '{count} user affinity', + affinityNoUsers: 'No affinity users', + affinityDetailTitle: 'Affinity Details', + affinityPinnedLabel: 'Pinned', expired: 'Expired', proxy: 'Proxy', noProxy: 'No Proxy', diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index 2f3eb2903f..d8a378b147 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -2319,7 +2319,7 @@ export default { // Quota control (Anthropic OAuth/SetupToken only) quotaControl: { title: '配额控制', - hint: '配置费用窗口、会话限制、客户端亲和等调度控制。', + hint: '配置亲和配置(用户亲和、客户端亲和)、费用窗口、会话限制等调度控制。', windowCost: { label: '5h窗口费用控制', hint: '限制账号在5小时窗口内的费用使用', @@ -2380,12 +2380,15 @@ export default { hint: '启用后,新会话会优先调度到该客户端之前使用过的账号,避免频繁切换账号' } }, + affinityConfigTitle: '亲和配置', + affinityConfigHint: '在配额控制下配置用户亲和与客户端亲和调度规则。', affinityNoClients: '无亲和客户端', - affinityClients: '{count} 个亲和客户端:', + affinityClients: '{count} 个客户端亲和', + affinityClientCountLabel: '{count} 个客户端', affinitySection: '客户端亲和', affinitySectionHint: '控制客户端在账号间的分布。通过配置区域阈值来平衡负载。', - affinityToggle: '启用客户端亲和', - affinityToggleHint: '新会话优先调度到该客户端之前使用过的账号', + affinityToggle: '启用亲和调度', + affinityToggleHint: '启用后,新会话会优先命中用户亲和与客户端亲和规则', affinityBase: '基础限额(绿区)', affinityBasePlaceholder: '留空表示不限制', affinityBaseHint: '绿区最大客户端数量(完整优先级调度)', @@ -2394,6 +2397,28 @@ export default { affinityBufferPlaceholder: '例如 3', affinityBufferHint: '黄区允许的额外客户端数量(降级优先级调度)', affinityBufferInfinite: '不限制', + affinityAllowSwitch: '允许切换', + affinityAllowSwitchHint: '亲和账号不可用时允许调度到其他账号', + affinityAllowSwitchWarning: '关闭后,亲和账号不可用将直接返回错误', + affinityUserSection: '用户亲和', + affinityUserSectionHint: '配置用户维度的亲和容量与调度规则', + affinityUserBase: '基础限额(绿区)', + affinityUserBaseHint: '绿区最大用户数量(完整优先级调度)', + affinityUserBuffer: '缓冲区(黄区)', + affinityUserBufferHint: '黄区允许的额外用户数量(降级优先级调度)', + affinityClientSection: '客户端亲和', + affinityClientSectionHint: '配置客户端维度的亲和容量与调度规则', + affinityPerUserLimit: '每用户客户端限制', + affinityPerUserLimitHint: '限制每个用户可使用的不同客户端数量', + affinityPerUserMax: '最大客户端数', + affinityPinnedUsers: '指定亲和用户', + affinityPinnedUsersHint: '预先绑定的用户将自动创建亲和缓存,占用用户亲和名额', + affinityPinnedUsersSearch: '搜索用户邮箱或用户名...', + affinityPinnedUsersEmpty: '未指定亲和用户', + affinityUsers: '{count} 个用户亲和', + affinityNoUsers: '无亲和用户', + affinityDetailTitle: '亲和详情', + affinityPinnedLabel: '指定', expired: '已过期', proxy: '代理', noProxy: '无代理', diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 469ac4afa1..7d8b12a1dc 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -731,6 +731,11 @@ export interface Account { affinity_client_count?: number | null affinity_clients?: string[] | null + // 二维亲和扩展 + affinity_allow_switch?: boolean | null // 允许切换 + affinity_user_count?: number | null // 关联用户数 + pinned_user_ids?: number[] | null // 指定亲和用户 ID 列表 + // API Key 账号配额限制 quota_limit?: number | null quota_used?: number | null @@ -755,6 +760,27 @@ export interface Account { current_rpm?: number | null // 当前分钟 RPM 计数 } +// Affinity Details types +export interface AffinityClientInfo { + client_id: string + last_active: string +} + +export interface AffinityUserGroup { + user_id: number + user_email: string + client_count: number + is_pinned: boolean + clients: AffinityClientInfo[] +} + +export interface AffinityDetailsResponse { + users: AffinityUserGroup[] + total_users: number + total_clients: number + pinned_users: number[] +} + // Account Usage types export interface WindowStats { requests: number