mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
fix: convert []byte to *string for JSONB columns in sora_task_repo
pq driver treats []byte as bytea, not jsonb. Convert charJSON and reqBody to *string before passing to ExecContext for character_info and request_body JSONB columns. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 账号
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 亲和客户端信息(含最后活跃时间)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 的空实现
|
||||
|
||||
Reference in New Issue
Block a user