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:
erio
2026-03-07 19:52:13 +08:00
co-authored by Claude Opus 4.6
parent a64dca7c41
commit e617dfa009
10 changed files with 140 additions and 11 deletions
@@ -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
+15 -11
View File
@@ -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()
+3
View File
@@ -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 的空实现