diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index f902101784..6e9ba868fc 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -8,6 +8,7 @@ import ( "encoding/json" "errors" "fmt" + "log/slog" "net/http" "strconv" "strings" @@ -619,6 +620,16 @@ 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 + } + } + // 确定是否跳过混合渠道检查 skipCheck := req.ConfirmMixedChannelRisk != nil && *req.ConfirmMixedChannelRisk @@ -655,6 +666,15 @@ 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) + } + response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account)) } @@ -1453,6 +1473,39 @@ func (h *AccountHandler) GetAffinityClients(c *gin.Context) { 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)) + result := make([]int64, 0, len(a)+len(b)) + for _, id := range a { + if _, ok := seen[id]; !ok { + seen[id] = struct{}{} + result = append(result, id) + } + } + for _, id := range b { + if _, ok := seen[id]; !ok { + seen[id] = struct{}{} + result = append(result, id) + } + } + return result +} + // GetTempUnschedulable handles getting temporary unschedulable status // GET /api/v1/admin/accounts/:id/temp-unschedulable func (h *AccountHandler) GetTempUnschedulable(c *gin.Context) { diff --git a/backend/internal/repository/gateway_cache.go b/backend/internal/repository/gateway_cache.go index 0d064c924a..a3cef7f092 100644 --- a/backend/internal/repository/gateway_cache.go +++ b/backend/internal/repository/gateway_cache.go @@ -28,12 +28,15 @@ var ( getAffinityClientsLua string //go:embed lua/get_affinity_clients_with_scores.lua getAffinityClientsWithScoresLua string + //go:embed lua/clear_account_affinity.lua + clearAccountAffinityLua string getAffinityScript = redis.NewScript(getAffinityLua) updateAffinityScript = redis.NewScript(updateAffinityLua) getAffinityCountScript = redis.NewScript(getAffinityCountLua) getAffinityClientsScript = redis.NewScript(getAffinityClientsLua) getAffinityClientsWithScoresScript = redis.NewScript(getAffinityClientsWithScoresLua) + clearAccountAffinityScript = redis.NewScript(clearAccountAffinityLua) ) type gatewayCache struct { @@ -269,3 +272,25 @@ func (c *gatewayCache) GetAccountAffinityClientsWithScores( return result, nil } + +// ClearAccountAffinity 清除指定账号在所有分组的亲和记录(正向+反向索引)。 +// 对每个 groupID 执行 Lua 脚本:读取反向索引获取所有客户端, +// 从每个客户端的正向索引中移除该账号,然后删除反向索引。 +func (c *gatewayCache) ClearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) error { + if len(groupIDs) == 0 { + return nil + } + + ensureScriptLoaded(ctx, c.rdb, clearAccountAffinityScript) + + pipe := c.rdb.Pipeline() + for _, gID := range groupIDs { + revKey := buildAffinityReverseKey(gID, accountID) + clearAccountAffinityScript.Run(ctx, pipe, []string{revKey}, gID, accountID) + } + _, err := pipe.Exec(ctx) + if err != nil && err != redis.Nil { + return err + } + return nil +} diff --git a/backend/internal/repository/lua/clear_account_affinity.lua b/backend/internal/repository/lua/clear_account_affinity.lua new file mode 100644 index 0000000000..e125be1690 --- /dev/null +++ b/backend/internal/repository/lua/clear_account_affinity.lua @@ -0,0 +1,29 @@ +-- 清除单个账号在指定分组的所有亲和记录(正向+反向) +-- KEYS[1] = client_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 + return 0 +end + +-- 从每个客户端的正向索引中移除该账号 +for _, client_id in ipairs(clients) do + local fwd_key = 'client_affinity:' .. group_id .. ':' .. client_id + redis.call('ZREM', fwd_key, account_id) + -- 如果正向索引为空,删除 key + if redis.call('ZCARD', fwd_key) == 0 then + redis.call('DEL', fwd_key) + end +end + +-- 删除反向索引 +redis.call('DEL', rev_key) + +return #clients diff --git a/backend/internal/repository/sora_task_repo.go b/backend/internal/repository/sora_task_repo.go index 49a101b139..822b61e41b 100644 --- a/backend/internal/repository/sora_task_repo.go +++ b/backend/internal/repository/sora_task_repo.go @@ -17,13 +17,15 @@ func NewSoraTaskRepository(sqlDB *sql.DB) service.SoraTaskRepository { } func (r *SoraTaskRepository) Create(ctx context.Context, task *service.SoraTask) error { - charJSON, _ := json.Marshal(task.CharacterInfo) - if task.CharacterInfo == nil { - charJSON = nil + var charStr, reqStr *string + if task.CharacterInfo != nil { + b, _ := json.Marshal(task.CharacterInfo) + s := string(b) + charStr = &s } - var reqBody []byte if len(task.RequestBody) > 0 { - reqBody = task.RequestBody + s := string(task.RequestBody) + reqStr = &s } _, err := r.db.ExecContext(ctx, ` @@ -36,8 +38,8 @@ func (r *SoraTaskRepository) Create(ctx context.Context, task *service.SoraTask) task.ID, task.AccountID, task.APIKeyID, task.UpstreamTaskID, task.ObjectType, task.Model, task.Prompt, task.Status, task.Progress, task.VideoURL, task.StoredKey, task.StorageType, - task.ShareID, charJSON, task.ErrorMessage, task.ErrorType, - reqBody, task.Seconds, task.Size, task.CreatedAt, task.CompletedAt, + task.ShareID, charStr, task.ErrorMessage, task.ErrorType, + reqStr, task.Seconds, task.Size, task.CreatedAt, task.CompletedAt, ) return err } @@ -56,9 +58,11 @@ func (r *SoraTaskRepository) GetByIDAndAPIKey(ctx context.Context, id string, ap } func (r *SoraTaskRepository) Update(ctx context.Context, task *service.SoraTask) error { - charJSON, _ := json.Marshal(task.CharacterInfo) - if task.CharacterInfo == nil { - charJSON = nil + var charStr *string + if task.CharacterInfo != nil { + b, _ := json.Marshal(task.CharacterInfo) + s := string(b) + charStr = &s } _, err := r.db.ExecContext(ctx, ` @@ -71,7 +75,7 @@ func (r *SoraTaskRepository) Update(ctx context.Context, task *service.SoraTask) WHERE id = $1`, task.ID, task.UpstreamTaskID, task.Status, task.Progress, task.VideoURL, task.StoredKey, task.StorageType, - task.ShareID, charJSON, + task.ShareID, charStr, task.ErrorMessage, task.ErrorType, task.CompletedAt, task.Seconds, task.Size, task.ObjectType, ) diff --git a/backend/internal/service/antigravity_smart_retry_test.go b/backend/internal/service/antigravity_smart_retry_test.go index 07d55031b0..2999172612 100644 --- a/backend/internal/service/antigravity_smart_retry_test.go +++ b/backend/internal/service/antigravity_smart_retry_test.go @@ -45,6 +45,9 @@ func (c *stubSmartRetryCache) GetAccountAffinityClientsBatch(_ context.Context, func (c *stubSmartRetryCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { return nil, nil } +func (c *stubSmartRetryCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} // mockSmartRetryUpstream 用于 handleSmartRetry 测试的 mock upstream type mockSmartRetryUpstream struct { diff --git a/backend/internal/service/gateway_affinity_scheduling_test.go b/backend/internal/service/gateway_affinity_scheduling_test.go index 7afe243036..0ad7bfb466 100644 --- a/backend/internal/service/gateway_affinity_scheduling_test.go +++ b/backend/internal/service/gateway_affinity_scheduling_test.go @@ -55,6 +55,9 @@ func (m *mockAffinityCache) GetAccountAffinityClientsBatch(_ context.Context, _ func (m *mockAffinityCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { return nil, nil } +func (m *mockAffinityCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} // --------------------------------------------------------------------------- // Helper: 构造启用了客户端亲和的 Anthropic 账号 diff --git a/backend/internal/service/gateway_hotpath_optimization_test.go b/backend/internal/service/gateway_hotpath_optimization_test.go index 4745b6248f..0203c3311f 100644 --- a/backend/internal/service/gateway_hotpath_optimization_test.go +++ b/backend/internal/service/gateway_hotpath_optimization_test.go @@ -158,6 +158,9 @@ func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityClientsBatch(_ context func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { return nil, nil } +func (s *stickyGatewayCacheHotpathStub) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} func (s *modelsListAccountRepoStub) ListSchedulableByGroupID(ctx context.Context, groupID int64) ([]Account, error) { s.listByGroupCalls.Add(1) diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 5dde09cfce..c3a53e0f4b 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -382,6 +382,9 @@ type GatewayCache interface { 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 } // AffinityClient 亲和客户端信息(含最后活跃时间) diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 3a67b93e96..ffa37d63fd 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -297,6 +297,9 @@ func (c *stubGatewayCache) GetAccountAffinityClientsBatch(_ context.Context, _ m func (c *stubGatewayCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) { return nil, nil } +func (c *stubGatewayCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} func TestOpenAISelectAccountWithLoadAwareness_FiltersUnschedulable(t *testing.T) { now := time.Now() diff --git a/backend/internal/testutil/stubs.go b/backend/internal/testutil/stubs.go index bae3e7d6f9..36375b66b7 100644 --- a/backend/internal/testutil/stubs.go +++ b/backend/internal/testutil/stubs.go @@ -112,6 +112,9 @@ func (c StubGatewayCache) GetAccountAffinityClientsBatch(_ context.Context, _ ma func (c StubGatewayCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]service.AffinityClient, error) { return nil, nil } +func (c StubGatewayCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error { + return nil +} // ============================================================ // StubSessionLimitCache — service.SessionLimitCache 的空实现