From 0fd2e9216d296b67054fab70d07c3bf36cada679 Mon Sep 17 00:00:00 2001 From: shaw Date: Mon, 6 Jul 2026 11:43:16 +0800 Subject: [PATCH] =?UTF-8?q?fix(scheduler):=20=E4=BF=AE=E5=A4=8D=20OpenAI?= =?UTF-8?q?=20=E9=AB=98=E7=BA=A7=E8=B0=83=E5=BA=A6=E5=99=A8=E5=AE=A1?= =?UTF-8?q?=E8=AE=A1=E5=8F=91=E7=8E=B0=E7=9A=84=E6=AD=A3=E7=A1=AE=E6=80=A7?= =?UTF-8?q?=E4=B8=8E=E6=80=A7=E8=83=BD=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 针对 #3692 合并后审计发现的问题集中修复: - previous_response_id 剥离条件改为按 call_id 全覆盖校验, 部分可重建的工具续链不再被误剥离(不受开关门控的行为回归) - 粘性加权回退路径补分组归属校验并清理失效绑定,杜绝跨分组账号泄漏 - 账号列表页:无 OpenAI 账号时跳过分数计算、过滤池限定 openai 平台、 负载批查合并为账号并集一次查询,消除全表扫描与 Redis N+1 - 订阅优先模式下常规池不可用时回退订阅池等待计划, busy-but-waitable 的订阅账号不再导致请求硬失败 - TopK/权重 DB 覆盖显式受总开关门控,与兄弟子开关语义一致 - 前端未分组 OpenAI 账号回退展示基础分,不再显示 "-" - ListAllWithFilters 等能力正式进入 AccountRepository/AdminService 接口, 移除匿名接口断言与静默降级;负载批查失败补 warn 日志 - SelectAccountWithSchedulerForCapability 增加显式 previousResponseCanMove 参数,移除 "previous_response_can_move" 魔法字符串哨兵 - 设置写入路径补"基础权重不得全为零"聚合校验; 运行时设置批量读取失败的降级路径覆盖全部键并留痕 --- .../internal/handler/admin/account_handler.go | 173 +++++++++----- backend/internal/handler/grok_media.go | 1 + .../handler/openai_chat_completions.go | 1 + backend/internal/handler/openai_embeddings.go | 1 + .../handler/openai_gateway_count_tokens.go | 1 + .../handler/openai_gateway_handler.go | 13 +- backend/internal/server/api_contract_test.go | 4 + backend/internal/service/account_service.go | 3 + .../service/account_service_delete_test.go | 4 + backend/internal/service/admin_service.go | 14 +- .../service/admin_service_bulk_update_test.go | 4 + .../service/admin_service_search_test.go | 4 + .../service/gateway_multiplatform_test.go | 3 + .../service/gemini_multiplatform_test.go | 3 + .../service/openai_account_scheduler.go | 66 ++++-- .../service/openai_account_scheduler_test.go | 145 +++++++++++- .../service/openai_tool_continuation.go | 78 ++++++ .../service/openai_tool_continuation_test.go | 106 +++++++++ .../service/ratelimit_session_window_test.go | 3 + backend/internal/service/setting_service.go | 30 ++- frontend/src/views/admin/AccountsView.vue | 13 +- .../AccountsView.schedulerScore.spec.ts | 224 ++++++++++++++++++ 22 files changed, 798 insertions(+), 96 deletions(-) create mode 100644 frontend/src/views/admin/__tests__/AccountsView.schedulerScore.spec.ts diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index efe95801c3..8c91245fbf 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -196,14 +196,6 @@ type AccountSchedulerGroupScore struct { const accountListGroupUngroupedQueryValue = "ungrouped" -type openAIAccountSchedulerScorePoolLister interface { - ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]service.Account, error) -} - -type accountSchedulerScoreFilterPoolLister interface { - ListAccountsForSchedulerScoreFilter(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, error) -} - func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, account *service.Account) AccountWithConcurrency { item := AccountWithConcurrency{ Account: dto.AccountFromService(account), @@ -250,33 +242,27 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac return item } -func (h *AccountHandler) scoreOpenAIAccountSchedulerPool(ctx context.Context, accounts []service.Account) map[int64]AccountSchedulerScore { +// scoreOpenAIAccountSchedulerPool 对池内 OpenAI 账号计算调度分数快照。 +// loadMap 为共享的账号负载数据(含池内全部账号即可,多余条目无害);传 nil 时自行批查。 +func (h *AccountHandler) scoreOpenAIAccountSchedulerPool(ctx context.Context, accounts []service.Account, loadMap map[int64]*service.AccountLoadInfo) map[int64]AccountSchedulerScore { if len(accounts) == 0 { return nil } openAIAccounts := make([]*service.Account, 0, len(accounts)) - loadReq := make([]service.AccountWithConcurrency, 0, len(accounts)) for i := range accounts { account := &accounts[i] if account.Platform != service.PlatformOpenAI { continue } openAIAccounts = append(openAIAccounts, account) - loadReq = append(loadReq, service.AccountWithConcurrency{ - ID: account.ID, - MaxConcurrency: account.EffectiveLoadFactor(), - }) } if len(openAIAccounts) == 0 { return nil } - loadMap := map[int64]*service.AccountLoadInfo{} - if h.concurrencyService != nil { - if batchLoad, err := h.concurrencyService.GetAccountsLoadBatch(ctx, loadReq); err == nil && batchLoad != nil { - loadMap = batchLoad - } + if loadMap == nil { + loadMap = h.fetchOpenAIAccountLoadMap(ctx, openAIAccounts) } var scores map[int64]service.OpenAIAccountSchedulerScoreSnapshot @@ -297,6 +283,36 @@ func (h *AccountHandler) scoreOpenAIAccountSchedulerPool(ctx context.Context, ac return result } +// fetchOpenAIAccountLoadMap 一次性批查给定 OpenAI 账号的负载数据; +// 失败时记录日志并返回空表(分数按零负载计算,属可接受降级)。 +func (h *AccountHandler) fetchOpenAIAccountLoadMap(ctx context.Context, openAIAccounts []*service.Account) map[int64]*service.AccountLoadInfo { + loadMap := map[int64]*service.AccountLoadInfo{} + if h.concurrencyService == nil || len(openAIAccounts) == 0 { + return loadMap + } + seen := make(map[int64]struct{}, len(openAIAccounts)) + loadReq := make([]service.AccountWithConcurrency, 0, len(openAIAccounts)) + for _, account := range openAIAccounts { + if account == nil { + continue + } + if _, ok := seen[account.ID]; ok { + continue + } + seen[account.ID] = struct{}{} + loadReq = append(loadReq, service.AccountWithConcurrency{ + ID: account.ID, + MaxConcurrency: account.EffectiveLoadFactor(), + }) + } + if batchLoad, err := h.concurrencyService.GetAccountsLoadBatch(ctx, loadReq); err != nil { + slog.Warn("openai_scheduler_score_load_batch_failed", "error", err) + } else if batchLoad != nil { + loadMap = batchLoad + } + return loadMap +} + func (h *AccountHandler) buildOpenAIAccountSchedulerScores( ctx context.Context, accounts []service.Account, @@ -309,12 +325,6 @@ func (h *AccountHandler) buildOpenAIAccountSchedulerScores( filterPool = accounts } - baseScores := make(map[int64]*AccountSchedulerScore) - for accountID, score := range h.scoreOpenAIAccountSchedulerPool(ctx, filterPool) { - copiedScore := score - baseScores[accountID] = &copiedScore - } - pageOpenAIAccountIDs := make(map[int64]struct{}) groupIDs := make(map[int64]struct{}) for i := range accounts { @@ -338,7 +348,48 @@ func (h *AccountHandler) buildOpenAIAccountSchedulerScores( } } if len(pageOpenAIAccountIDs) == 0 { - return baseScores, nil + return nil, nil + } + + // 先取各分组池,再对"过滤池 ∪ 分组池"的账号并集做一次负载批查, + // 避免每个池各查一次 Redis 的 N+1。 + groupIDList := make([]int64, 0, len(groupIDs)) + for groupID := range groupIDs { + groupIDList = append(groupIDList, groupID) + } + sort.Slice(groupIDList, func(i, j int) bool { return groupIDList[i] < groupIDList[j] }) + + groupPools := make(map[int64][]service.Account, len(groupIDList)) + if h.adminService != nil { + for _, groupID := range groupIDList { + gid := groupID + pool, err := h.adminService.ListOpenAISchedulableAccountsForSchedulerScore(ctx, &gid) + if err != nil { + slog.Warn("openai_scheduler_group_score_pool_failed", "group_id", gid, "error", err) + continue + } + groupPools[gid] = pool + } + } + + loadUnion := make([]*service.Account, 0, len(filterPool)) + collectOpenAIAccounts := func(pool []service.Account) { + for i := range pool { + if pool[i].Platform == service.PlatformOpenAI { + loadUnion = append(loadUnion, &pool[i]) + } + } + } + collectOpenAIAccounts(filterPool) + for _, pool := range groupPools { + collectOpenAIAccounts(pool) + } + loadMap := h.fetchOpenAIAccountLoadMap(ctx, loadUnion) + + baseScores := make(map[int64]*AccountSchedulerScore) + for accountID, score := range h.scoreOpenAIAccountSchedulerPool(ctx, filterPool, loadMap) { + copiedScore := score + baseScores[accountID] = &copiedScore } groupScoresByAccount := make(map[int64][]AccountSchedulerGroupScore) @@ -346,7 +397,7 @@ func (h *AccountHandler) buildOpenAIAccountSchedulerScores( if len(pool) == 0 { return } - scores := h.scoreOpenAIAccountSchedulerPool(ctx, pool) + scores := h.scoreOpenAIAccountSchedulerPool(ctx, pool, loadMap) for accountID, schedulerScore := range scores { if _, ok := pageOpenAIAccountIDs[accountID]; !ok { continue @@ -365,37 +416,27 @@ func (h *AccountHandler) buildOpenAIAccountSchedulerScores( } } - if lister, ok := h.adminService.(openAIAccountSchedulerScorePoolLister); ok { - groupIDList := make([]int64, 0, len(groupIDs)) - for groupID := range groupIDs { - groupIDList = append(groupIDList, groupID) + for _, groupID := range groupIDList { + gid := groupID + pool, ok := groupPools[gid] + if !ok { + continue } - sort.Slice(groupIDList, func(i, j int) bool { return groupIDList[i] < groupIDList[j] }) - - for _, groupID := range groupIDList { - gid := groupID - pool, err := lister.ListOpenAISchedulableAccountsForSchedulerScore(ctx, &gid) - if err != nil { - slog.Warn("openai_scheduler_group_score_pool_failed", "group_id", gid, "error", err) - continue - } - groupNameByID := make(map[int64]string) - groupPriorityByAccount := make(map[int64]int) - for i := range pool { - account := &pool[i] - for _, accountGroup := range account.AccountGroups { - if accountGroup.GroupID != gid { - continue - } - groupPriorityByAccount[account.ID] = accountGroup.Priority - if accountGroup.Group != nil { - groupNameByID[gid] = accountGroup.Group.Name - } + groupNameByID := make(map[int64]string) + groupPriorityByAccount := make(map[int64]int) + for i := range pool { + account := &pool[i] + for _, accountGroup := range account.AccountGroups { + if accountGroup.GroupID != gid { + continue + } + groupPriorityByAccount[account.ID] = accountGroup.Priority + if accountGroup.Group != nil { + groupNameByID[gid] = accountGroup.Group.Name } } - scoreGroupPool(&gid, groupNameByID, groupPriorityByAccount, pool) } - + scoreGroupPool(&gid, groupNameByID, groupPriorityByAccount, pool) } for accountID := range groupScoresByAccount { @@ -417,11 +458,9 @@ func (h *AccountHandler) listAccountSchedulerScoreFilterPool( if h.adminService == nil || (platform != "" && platform != service.PlatformOpenAI) { return nil } - lister, ok := h.adminService.(accountSchedulerScoreFilterPoolLister) - if !ok { - return nil - } - accounts, err := lister.ListAccountsForSchedulerScoreFilter(ctx, platform, accountType, status, search, groupID, privacyMode) + // 池只用于 OpenAI 分数计算(非 OpenAI 账号会在打分时被丢弃), + // 无论列表页平台过滤为何,查询一律限定 openai,避免无过滤时全表扫描。 + accounts, err := h.adminService.ListAccountsForSchedulerScoreFilter(ctx, service.PlatformOpenAI, accountType, status, search, groupID, privacyMode) if err != nil { slog.Warn("openai_scheduler_filter_score_pool_failed", "error", err) return nil @@ -481,8 +520,20 @@ func (h *AccountHandler) List(c *gin.Context) { var windowCosts map[int64]float64 var activeSessions map[int64]int var rpmCounts map[int64]int - schedulerFilterPool := h.listAccountSchedulerScoreFilterPool(c.Request.Context(), platform, accountType, status, search, groupID, privacyMode) - schedulerScores, schedulerGroupScores := h.buildOpenAIAccountSchedulerScores(c.Request.Context(), accounts, schedulerFilterPool) + // 仅当前页存在 OpenAI 账号时才计算调度分数,避免为空结果付出池查询开销。 + var schedulerScores map[int64]*AccountSchedulerScore + var schedulerGroupScores map[int64][]AccountSchedulerGroupScore + pageHasOpenAIAccounts := false + for i := range accounts { + if accounts[i].Platform == service.PlatformOpenAI { + pageHasOpenAIAccounts = true + break + } + } + if pageHasOpenAIAccounts { + schedulerFilterPool := h.listAccountSchedulerScoreFilterPool(c.Request.Context(), platform, accountType, status, search, groupID, privacyMode) + schedulerScores, schedulerGroupScores = h.buildOpenAIAccountSchedulerScores(c.Request.Context(), accounts, schedulerFilterPool) + } // 始终获取并发数(Redis ZCARD,极低开销) if h.concurrencyService != nil { diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index 8e236ea49f..4fd1411b23 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -174,6 +174,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. service.OpenAIUpstreamTransportHTTPSSE, "", false, + false, service.PlatformGrok, ) if err != nil { diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index ca43a2cff3..baff1dcbd6 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -145,6 +145,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, requestPlatform, ) if err != nil { diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index a80c7f7d96..8be533c723 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -117,6 +117,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { service.OpenAIUpstreamTransportHTTPSSE, service.OpenAIEndpointCapabilityEmbeddings, false, + false, ) if err != nil { reqLog.Warn("openai_embeddings.account_select_failed", diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index ec530e8ab6..fc9c4d5df7 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -110,6 +110,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, openAICompatibleRequestPlatform(apiKey), ) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index ccafde7b02..e177edb43d 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -350,6 +350,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, requireCompact, + false, requestPlatform, ) if err != nil { @@ -783,6 +784,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, false, + false, requestPlatform, ) if err != nil { @@ -1266,8 +1268,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "previous_response_id must be a response.id (resp_*), not a message id") return } - firstMessageToolContext := service.ValidateFunctionCallOutputContextBytes(firstMessage) - previousResponseCanMove := !firstMessageToolContext.HasFunctionCallOutput || firstMessageToolContext.HasToolCallContext + firstMessageToolCoverage := service.AnalyzeToolCallOutputContextCoverageBytes(firstMessage) + previousResponseCanMove := !firstMessageToolCoverage.HasFunctionCallOutput || firstMessageToolCoverage.ContextCoversAllCallIDs reqLog = reqLog.With( zap.Bool("ws_ingress", true), zap.String("model", reqModel), @@ -1383,13 +1385,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { requiredTransport, service.OpenAIEndpointCapabilityChatCompletions, false, + previousResponseCanMove, requestPlatform, - func() string { - if previousResponseCanMove { - return "previous_response_can_move" - } - return "" - }(), ) if err != nil { reqLog.Warn("openai.websocket_account_select_failed", diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index dfa480dd18..9b3f2dcdd1 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -1736,6 +1736,10 @@ func (s *stubAccountRepo) List(ctx context.Context, params pagination.Pagination return nil, nil, errors.New("not implemented") } +func (s *stubAccountRepo) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]service.Account, error) { + return nil, nil +} + func (s *stubAccountRepo) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, *pagination.PaginationResult, error) { return nil, nil, errors.New("not implemented") } diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index dcba614c2c..5956684f98 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -39,6 +39,9 @@ type AccountRepository interface { List(ctx context.Context, params pagination.PaginationParams) ([]Account, *pagination.PaginationResult, error) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) + // ListAllWithFilters 返回符合过滤条件的全部账号(不分页),用于账号列表页 + // 计算 OpenAI 调度分数的过滤范围池。 + ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) ListActive(ctx context.Context) ([]Account, error) ListOAuthRefreshCandidates(ctx context.Context) ([]Account, error) diff --git a/backend/internal/service/account_service_delete_test.go b/backend/internal/service/account_service_delete_test.go index a304356c09..ee6163239e 100644 --- a/backend/internal/service/account_service_delete_test.go +++ b/backend/internal/service/account_service_delete_test.go @@ -79,6 +79,10 @@ func (s *accountRepoStub) List(ctx context.Context, params pagination.Pagination panic("unexpected List call") } +func (s *accountRepoStub) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStub) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { panic("unexpected ListWithFilters call") } diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index ce59c34475..ebf1e7e404 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -78,6 +78,12 @@ type AdminService interface { // Account management ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string, sortBy, sortOrder string) ([]Account, int64, error) + // ListAccountsForSchedulerScoreFilter 返回符合过滤条件的全部账号(不分页), + // 作为账号列表页计算 OpenAI 调度分数的过滤范围池。 + ListAccountsForSchedulerScoreFilter(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) + // ListOpenAISchedulableAccountsForSchedulerScore 返回指定分组(nil 为未分组)内 + // 可调度的 OpenAI 账号,用于按组计算调度分数。 + ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]Account, error) GetAccount(ctx context.Context, id int64) (*Account, error) GetAccountsByIDs(ctx context.Context, ids []int64) ([]*Account, error) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) @@ -2622,13 +2628,7 @@ func (s *adminServiceImpl) ListAccountsForSchedulerScoreFilter(ctx context.Conte if s == nil || s.accountRepo == nil { return nil, nil } - lister, ok := s.accountRepo.(interface { - ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) - }) - if !ok { - return nil, nil - } - return lister.ListAllWithFilters(ctx, platform, accountType, status, search, groupID, privacyMode) + return s.accountRepo.ListAllWithFilters(ctx, platform, accountType, status, search, groupID, privacyMode) } func (s *adminServiceImpl) ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]Account, error) { diff --git a/backend/internal/service/admin_service_bulk_update_test.go b/backend/internal/service/admin_service_bulk_update_test.go index df415295b1..2f44b1741d 100644 --- a/backend/internal/service/admin_service_bulk_update_test.go +++ b/backend/internal/service/admin_service_bulk_update_test.go @@ -88,6 +88,10 @@ func (s *accountRepoStubForBulkUpdate) ListByGroup(_ context.Context, groupID in return nil, nil } +func (s *accountRepoStubForBulkUpdate) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStubForBulkUpdate) ListWithFilters(_ context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { s.listCalled = true s.lastListParams = params diff --git a/backend/internal/service/admin_service_search_test.go b/backend/internal/service/admin_service_search_test.go index 595e99e344..76acd1b5e2 100644 --- a/backend/internal/service/admin_service_search_test.go +++ b/backend/internal/service/admin_service_search_test.go @@ -25,6 +25,10 @@ type accountRepoStubForAdminList struct { listWithFiltersErr error } +func (s *accountRepoStubForAdminList) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + return nil, nil +} + func (s *accountRepoStubForAdminList) ListWithFilters(_ context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { s.listWithFiltersCalls++ s.listWithFiltersParams = params diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go index f843ba3e45..35a60c8124 100644 --- a/backend/internal/service/gateway_multiplatform_test.go +++ b/backend/internal/service/gateway_multiplatform_test.go @@ -95,6 +95,9 @@ func (m *mockAccountRepoForPlatform) List(ctx context.Context, params pagination func (m *mockAccountRepoForPlatform) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { return nil, nil, nil } +func (m *mockAccountRepoForPlatform) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) { + return nil, nil +} func (m *mockAccountRepoForPlatform) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) { return nil, nil } diff --git a/backend/internal/service/gemini_multiplatform_test.go b/backend/internal/service/gemini_multiplatform_test.go index c021e88edf..7d5ed0ec9e 100644 --- a/backend/internal/service/gemini_multiplatform_test.go +++ b/backend/internal/service/gemini_multiplatform_test.go @@ -82,6 +82,9 @@ func (m *mockAccountRepoForGemini) List(ctx context.Context, params pagination.P func (m *mockAccountRepoForGemini) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) { return nil, nil, nil } +func (m *mockAccountRepoForGemini) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) { + return nil, nil +} func (m *mockAccountRepoForGemini) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) { return nil, nil } diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index dd65163abc..a0a4fbff3e 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -1037,6 +1037,14 @@ func (s *defaultOpenAIAccountScheduler) tryFallbackToWeightedSticky( if account == nil || !s.isAccountRequestCompatible(ctx, account, req) || !s.isAccountTransportCompatible(account, req.RequiredTransport) { continue } + // 粘性绑定只证明绑定时账号在分组内;账号被移出分组后绑定仍会在 TTL 内存活, + // 必须与 selectBySessionHash 一样重验分组归属,否则会把分组流量泄漏到组外账号。 + if !openAIStickyAccountMatchesGroup(account, req.GroupID) { + if accountID == req.StickyAccountID && strings.TrimSpace(req.SessionHash) != "" { + _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, req.SessionHash) + } + continue + } if req.RequireCompact && openAICompactSupportTier(account) == 0 { continue } @@ -1145,13 +1153,29 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( } if len(regularAccounts) > 0 { regularAttempt := s.trySelectByLoadBalancePool(ctx, req, regularAccounts, loadMap) - if regularAttempt.err != nil { + if regularAttempt.err != nil && !regularAttempt.noCompactCandidates { return nil, regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew, regularAttempt.err } if regularAttempt.result != nil { return regularAttempt.result, regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew, nil } - return s.finishLoadBalanceSelectionFallback(ctx, req, regularAttempt) + var result *AccountSelectionResult + candidateCount, topK, loadSkew := regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew + fallbackErr := regularAttempt.err + if regularAttempt.err == nil { + result, candidateCount, topK, loadSkew, fallbackErr = s.finishLoadBalanceSelectionFallback(ctx, req, regularAttempt) + if fallbackErr == nil && result != nil { + return result, candidateCount, topK, loadSkew, nil + } + } + // 常规池既无法获取也无法排队(含仅剩不支持 compact 的候选)时, + // 回退到订阅池的等待计划:busy-but-waitable 的订阅账号不应因常规池存在 + // 而被丢弃,否则开启订阅优先反而让本可排队成功的请求硬失败。 + subResult, subCandidateCount, subTopK, subLoadSkew, subErr := s.finishLoadBalanceSelectionFallback(ctx, req, attempt) + if subErr == nil && subResult != nil { + return subResult, subCandidateCount, subTopK, subLoadSkew, nil + } + return result, candidateCount, topK, loadSkew, fallbackErr } return s.finishLoadBalanceSelectionFallback(ctx, req, attempt) } @@ -1464,15 +1488,20 @@ func (s *OpenAIGatewayService) openAIAdvancedSchedulerRuntimeSettings(ctx contex lbTopKOverride = parsePositiveIntOverride(values[SettingKeyOpenAIAdvancedSchedulerLBTopK]) weightOverrides = parseOpenAIAdvancedSchedulerWeightOverrides(values) } else { - if value, err := repo.GetValue(dbCtx, openAIAdvancedSchedulerSettingKey); err == nil { - enabled = strings.EqualFold(strings.TrimSpace(value), "true") - } - if value, err := repo.GetValue(dbCtx, SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled); err == nil { - stickyWeightedEnabled = strings.EqualFold(strings.TrimSpace(value), "true") - } - if value, err := repo.GetValue(dbCtx, SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled); err == nil { - subscriptionPriorityEnabled = strings.EqualFold(strings.TrimSpace(value), "true") + // 批量读取失败时逐键降级,覆盖全部键(含 TopK/权重),避免只加载布尔开关 + // 而静默丢弃管理员配置的覆盖值;降级状态会被缓存一个 TTL,必须留痕。 + slog.Warn("openai_advanced_scheduler_settings_batch_load_failed", "error", err) + fallbackValues := make(map[string]string) + for _, key := range openAIAdvancedSchedulerRuntimeSettingKeys() { + if value, valueErr := repo.GetValue(dbCtx, key); valueErr == nil { + fallbackValues[key] = value + } } + enabled = strings.EqualFold(strings.TrimSpace(fallbackValues[openAIAdvancedSchedulerSettingKey]), "true") + stickyWeightedEnabled = strings.EqualFold(strings.TrimSpace(fallbackValues[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled]), "true") + subscriptionPriorityEnabled = strings.EqualFold(strings.TrimSpace(fallbackValues[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled]), "true") + lbTopKOverride = parsePositiveIntOverride(fallbackValues[SettingKeyOpenAIAdvancedSchedulerLBTopK]) + weightOverrides = parseOpenAIAdvancedSchedulerWeightOverrides(fallbackValues) } } @@ -1618,6 +1647,9 @@ func (s *OpenAIGatewayService) SelectAccountWithScheduler( return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact, PlatformOpenAI, false) } +// SelectAccountWithSchedulerForCapability 按能力要求调度账号。 +// previousResponseCanMove 表示首包 input 可自行重建工具续链,previous_response_id 允许跨账号迁移 +// (粘性加权模式下改为加权偏好而非硬粘连)。 func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( ctx context.Context, groupID *int64, @@ -1628,16 +1660,13 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( requiredTransport OpenAIUpstreamTransport, requiredCapability OpenAIEndpointCapability, requireCompact bool, + previousResponseCanMove bool, platformOverride ...string, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { platform := PlatformOpenAI - previousResponseCanMove := false if len(platformOverride) > 0 { platform = platformOverride[0] } - if len(platformOverride) > 1 { - previousResponseCanMove = strings.EqualFold(platformOverride[1], "previous_response_can_move") - } return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact, platform, previousResponseCanMove) } @@ -1853,6 +1882,11 @@ func (s *OpenAIGatewayService) openAIWSLBTopK() int { func (s *OpenAIGatewayService) openAIWSLBTopKForRequest(ctx context.Context) int { base := s.openAIWSLBTopK() settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx) + // DB 覆盖值与 stickyWeighted/subscriptionPriority 一样受总开关门控: + // 关闭高级调度器后所有调用方(含管理页分数快照)都应回到配置/默认行为。 + if !settings.enabled { + return base + } if settings.lbTopKOverride > 0 { return settings.lbTopKOverride } @@ -1920,6 +1954,10 @@ func (s *OpenAIGatewayService) openAIWSSchedulerWeights() GatewayOpenAIWSSchedul func (s *OpenAIGatewayService) openAIWSSchedulerWeightsForRequest(ctx context.Context) GatewayOpenAIWSSchedulerScoreWeightsView { weights := s.openAIWSSchedulerWeights() settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx) + // 同 openAIWSLBTopKForRequest:总开关关闭时不应用 DB 覆盖值。 + if !settings.enabled { + return weights + } return applyOpenAIAdvancedSchedulerWeightOverrides(weights, settings.weightOverrides) } diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index a61a923053..2ff7c25e5d 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -529,6 +529,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Embeddi OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -572,6 +573,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_AllowsG OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityChatCompletions, false, + false, PlatformGrok, ) require.NoError(t, err) @@ -774,6 +776,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedPreviousR OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityChatCompletions, false, + false, PlatformOpenAI, ) require.NoError(t, err) @@ -796,8 +799,8 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedPreviousR OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityChatCompletions, false, + true, PlatformOpenAI, - "previous_response_can_move", ) require.NoError(t, err) require.NotNil(t, selection) @@ -938,6 +941,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -1011,6 +1015,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityEmbeddings, false, + false, ) require.NoError(t, err) require.NotNil(t, selection) @@ -2935,3 +2940,141 @@ func TestDefaultOpenAIAccountScheduler_IsAccountTransportCompatible_Branches(t * func int64PtrForTest(v int64) *int64 { return &v } + +func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedFallbackSkipsOutOfGroupStickyAccount(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101081) + otherGroupID := int64(101082) + accounts := []Account{ + { + ID: 38001, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 10, + GroupIDs: []int64{groupID}, + }, + { + // 会话粘连绑定指向的账号已被移出请求分组(绑定 TTL 内账号改组的场景)。 + ID: 38002, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{otherGroupID}, + }, + } + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.LBTopK = 2 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky = 3 + cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{ + "openai:session_weighted_out_of_group": 38002, + }} + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{38001: false, 38002: true}, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerGroupAwareOpenAIAccountRepo{schedulerTestOpenAIAccountRepo{accounts: accounts}}, + cache: cache, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "session_weighted_out_of_group", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + false, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + // 组内唯一候选 38001 满并发:必须返回其等待计划,绝不能把请求泄漏到组外的粘连账号 38002。 + require.Equal(t, int64(38001), selection.Account.ID) + require.False(t, selection.Acquired) + require.NotNil(t, selection.WaitPlan) + require.Equal(t, int64(38001), selection.WaitPlan.AccountID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + // 失效的粘连绑定应被清理,避免后续请求反复走同一条泄漏路径。 + require.Positive(t, cache.deletedSessions["openai:session_weighted_out_of_group"]) +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityWaitsOnBusySubscriptionWhenRegularUnusable(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(101091) + accounts := []Account{ + { + // 订阅账号:支持 compact,但并发已满(busy-but-waitable)。 + ID: 38011, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + GroupIDs: []int64{groupID}, + Credentials: map[string]any{"plan_type": "team"}, + Extra: map[string]any{"openai_compact_supported": true}, + }, + { + // 常规账号:明确不支持 compact,无法服务本次请求。 + ID: 38012, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 9, + GroupIDs: []int64{groupID}, + Extra: map[string]any{"openai_compact_supported": false}, + }, + } + concurrencyCache := schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{38011: false, 38012: true}, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: newSchedulerTestSubscriptionPriorityConfig(), + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "session_subscription_wait", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + true, + ) + // 常规池无可用候选时,忙碌的订阅账号应产生等待计划,而不是直接返回 no available accounts。 + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(38011), selection.Account.ID) + require.False(t, selection.Acquired) + require.NotNil(t, selection.WaitPlan) + require.Equal(t, int64(38011), selection.WaitPlan.AccountID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) +} diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go index 6515c0c4e5..a507213701 100644 --- a/backend/internal/service/openai_tool_continuation.go +++ b/backend/internal/service/openai_tool_continuation.go @@ -215,6 +215,84 @@ func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValid return result } +// ToolCallOutputContextCoverage 描述 input 中工具输出与可重建上下文的覆盖关系, +// 用于判断剥离 previous_response_id 后上游能否仅凭 input 重建工具续链。 +type ToolCallOutputContextCoverage struct { + HasFunctionCallOutput bool + // ContextCoversAllCallIDs 表示每个工具输出的 call_id 都能在 input 内找到 + // 同 call_id 的工具调用上下文项或同 id 的 item_reference,且不存在缺失 call_id 的输出。 + // 任一输出无法由 input 自身重建时为 false,此时剥离 previous_response_id 会导致 + // 上游以 "No tool call found for function call output" 拒绝请求。 + ContextCoversAllCallIDs bool +} + +// AnalyzeToolCallOutputContextCoverageBytes 全量扫描 input,按 call_id 精确匹配工具输出 +// 与可重建上下文。不能复用 ValidateFunctionCallOutputContextBytes 的 HasToolCallContext: +// 该标志只代表"存在某一个上下文项",部分覆盖的续链仍会被上游拒绝。 +func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContextCoverage { + coverage := ToolCallOutputContextCoverage{} + if len(body) == 0 { + return coverage + } + input := parseRawJSONView(body).Get("input") + if !input.IsArray() { + return coverage + } + + missingCallID := false + var outputCallIDs map[string]struct{} + var contextIDs map[string]struct{} + input.ForEach(func(_, item gjson.Result) bool { + if !item.IsObject() { + return true + } + itemType := item.Get("type").String() + switch { + case isCodexToolCallOutputItemType(itemType): + coverage.HasFunctionCallOutput = true + callID := strings.TrimSpace(item.Get("call_id").String()) + if callID == "" { + missingCallID = true + return true + } + if outputCallIDs == nil { + outputCallIDs = make(map[string]struct{}) + } + outputCallIDs[callID] = struct{}{} + case isCodexToolCallContextItemType(itemType): + callID := strings.TrimSpace(item.Get("call_id").String()) + if callID == "" { + return true + } + if contextIDs == nil { + contextIDs = make(map[string]struct{}) + } + contextIDs[callID] = struct{}{} + case itemType == "item_reference": + idValue := strings.TrimSpace(item.Get("id").String()) + if idValue == "" { + return true + } + if contextIDs == nil { + contextIDs = make(map[string]struct{}) + } + contextIDs[idValue] = struct{}{} + } + return true + }) + + if !coverage.HasFunctionCallOutput || missingCallID { + return coverage + } + for callID := range outputCallIDs { + if _, ok := contextIDs[callID]; !ok { + return coverage + } + } + coverage.ContextCoversAllCallIDs = true + return coverage +} + // ValidateFunctionCallOutputContext 为 handler 提供低开销校验结果: // 1) 无工具输出直接返回 // 2) 若已存在工具调用上下文则提前返回 diff --git a/backend/internal/service/openai_tool_continuation_test.go b/backend/internal/service/openai_tool_continuation_test.go index 4610652b6c..569d89eff0 100644 --- a/backend/internal/service/openai_tool_continuation_test.go +++ b/backend/internal/service/openai_tool_continuation_test.go @@ -184,3 +184,109 @@ func TestValidateFunctionCallOutputContextBytesMatchesMapValidation(t *testing.T }) } } + +func TestAnalyzeToolCallOutputContextCoverageBytes(t *testing.T) { + cases := []struct { + name string + body map[string]any + hasOutput bool + coversAllIDs bool + }{ + { + name: "no_input", + body: map[string]any{"model": "gpt-5.1"}, + hasOutput: false, + coversAllIDs: false, + }, + { + name: "no_tool_output", + body: map[string]any{"input": []any{ + map[string]any{"type": "message", "content": "hi"}, + }}, + hasOutput: false, + coversAllIDs: false, + }, + { + name: "all_outputs_covered_by_context", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + name: "all_outputs_covered_by_item_reference", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "item_reference", "id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + // 关键回归用例:input 内存在某一个上下文项,但另一个输出的 call_id + // 只能由上游会话链(previous_response_id)解析——不可剥离。 + name: "partial_coverage_not_movable", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "unrelated_context_does_not_cover", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_x"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "output_missing_call_id_not_movable", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + }}, + hasOutput: true, + coversAllIDs: false, + }, + { + name: "mixed_context_and_reference_cover_all", + body: map[string]any{"input": []any{ + map[string]any{"type": "function_call", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_a"}, + map[string]any{"type": "function_call_output", "call_id": "call_b"}, + map[string]any{"type": "item_reference", "id": "call_b"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + { + name: "all_codex_output_types_covered", + body: map[string]any{"input": []any{ + map[string]any{"type": "tool_search_output", "call_id": "call_s"}, + map[string]any{"type": "tool_search_call", "call_id": "call_s"}, + map[string]any{"type": "mcp_tool_call_output", "call_id": "call_m"}, + map[string]any{"type": "mcp_tool_call", "call_id": "call_m"}, + }}, + hasOutput: true, + coversAllIDs: true, + }, + } + + for _, tt := range cases { + t.Run(tt.name, func(t *testing.T) { + bodyBytes, err := json.Marshal(tt.body) + require.NoError(t, err) + + coverage := AnalyzeToolCallOutputContextCoverageBytes(bodyBytes) + require.Equal(t, tt.hasOutput, coverage.HasFunctionCallOutput, "HasFunctionCallOutput") + require.Equal(t, tt.coversAllIDs, coverage.ContextCoversAllCallIDs, "ContextCoversAllCallIDs") + }) + } +} diff --git a/backend/internal/service/ratelimit_session_window_test.go b/backend/internal/service/ratelimit_session_window_test.go index cb19227e54..279a31ccdc 100644 --- a/backend/internal/service/ratelimit_session_window_test.go +++ b/backend/internal/service/ratelimit_session_window_test.go @@ -87,6 +87,9 @@ func (m *sessionWindowMockRepo) List(context.Context, pagination.PaginationParam func (m *sessionWindowMockRepo) ListWithFilters(context.Context, pagination.PaginationParams, string, string, string, string, int64, string) ([]Account, *pagination.PaginationResult, error) { panic("unexpected") } +func (m *sessionWindowMockRepo) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) { + panic("unexpected") +} func (m *sessionWindowMockRepo) ListByGroup(context.Context, int64) ([]Account, error) { panic("unexpected") } diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index 1024243bea..3aeb611418 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -1921,7 +1921,7 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting if err != nil { return nil, err } - if err := normalizeOpenAIAdvancedSchedulerOverrides(settings); err != nil { + if err := s.normalizeOpenAIAdvancedSchedulerOverrides(settings); err != nil { return nil, err } settings.PaymentVisibleMethodAlipaySource = alipaySource @@ -3940,7 +3940,7 @@ func formatOpenAIAdvancedSchedulerFloat(value float64) string { return strconv.FormatFloat(value, 'f', -1, 64) } -func normalizeOpenAIAdvancedSchedulerOverrides(settings *SystemSettings) error { +func (s *SettingService) normalizeOpenAIAdvancedSchedulerOverrides(settings *SystemSettings) error { lbTopK, err := normalizeOptionalPositiveIntString(settings.OpenAIAdvancedSchedulerLBTopK) if err != nil { return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_LB_TOP_K", "openai advanced scheduler TopK must be a positive integer or empty") @@ -3965,9 +3965,35 @@ func normalizeOpenAIAdvancedSchedulerOverrides(settings *SystemSettings) error { } *target = normalized } + + // 与 config.Validate 的 "scheduler_score_weights must not all be zero" 保持一致: + // 覆盖值(空则回退到生效的配置值)叠加后的基础权重和不允许为 0, + // 否则调度会静默退化为 TopK 内均匀随机。 + effective := s.openAIAdvancedSchedulerEffectiveWeights() + baseSum := resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightPriority, effective.Priority) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightLoad, effective.Load) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQueue, effective.Queue) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightErrorRate, effective.ErrorRate) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightTTFT, effective.TTFT) + + resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, effective.QuotaHeadroom) + if baseSum <= 0 { + return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_WEIGHT", "openai advanced scheduler base weights must not all be zero") + } return nil } +// resolveOpenAIAdvancedSchedulerWeight 返回覆盖值(已归一化的非空字符串),空则回退默认值。 +func resolveOpenAIAdvancedSchedulerWeight(normalized string, fallback float64) float64 { + if normalized == "" { + return fallback + } + value, err := strconv.ParseFloat(normalized, 64) + if err != nil { + return fallback + } + return value +} + func normalizeOptionalPositiveIntString(raw string) (string, error) { raw = strings.TrimSpace(raw) if raw == "" { diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index bd8fd6067a..b4e1630a85 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -679,14 +679,21 @@ const formatStickySchedulerScore = (score: AccountSchedulerGroupScore): string = } const getSchedulerScoreRows = (account: Account): AccountSchedulerGroupScore[] => { - if (!Array.isArray(account.scheduler_scores)) return [] - return account.scheduler_scores.filter(score => score.group_id != null) + const groupRows = Array.isArray(account.scheduler_scores) + ? account.scheduler_scores.filter(score => score.group_id != null) + : [] + if (groupRows.length) return groupRows + // 未分组账号没有分组维度分数,回退展示后端返回的基础分 + if (account.scheduler_score) { + return [{ group_id: null, ...account.scheduler_score }] + } + return [] } const formatSchedulerScoreGroup = (score: AccountSchedulerGroupScore): string => { if ('group_name' in score && score.group_name) return score.group_name if ('group_id' in score && score.group_id != null) return `#${score.group_id}` - return '-' + return t('admin.accounts.schedulerScore.ungrouped') } const loadSavedColumns = () => { diff --git a/frontend/src/views/admin/__tests__/AccountsView.schedulerScore.spec.ts b/frontend/src/views/admin/__tests__/AccountsView.schedulerScore.spec.ts new file mode 100644 index 0000000000..0865a6ec91 --- /dev/null +++ b/frontend/src/views/admin/__tests__/AccountsView.schedulerScore.spec.ts @@ -0,0 +1,224 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' +import { flushPromises, mount } from '@vue/test-utils' + +import AccountsView from '../AccountsView.vue' + +const { + listAccounts, + listWithEtag, + getBatchTodayStats, + getAllProxies, + getAllGroups +} = vi.hoisted(() => ({ + listAccounts: vi.fn(), + listWithEtag: vi.fn(), + getBatchTodayStats: vi.fn(), + getAllProxies: vi.fn(), + getAllGroups: vi.fn() +})) + +vi.mock('@/api/admin', () => ({ + adminAPI: { + accounts: { + list: listAccounts, + listWithEtag, + getBatchTodayStats, + delete: vi.fn(), + batchClearError: vi.fn(), + batchRefresh: vi.fn(), + toggleSchedulable: vi.fn() + }, + proxies: { + getAll: getAllProxies + }, + groups: { + getAll: getAllGroups + } + } +})) + +vi.mock('@/stores/app', () => ({ + useAppStore: () => ({ + showError: vi.fn(), + showSuccess: vi.fn(), + showInfo: vi.fn() + }) +})) + +vi.mock('@/stores/auth', () => ({ + useAuthStore: () => ({ + token: 'test-token' + }) +})) + +vi.mock('vue-i18n', async () => { + const actual = await vi.importActual('vue-i18n') + return { + ...actual, + useI18n: () => ({ + t: (key: string) => key + }) + } +}) + +// Render the scheduler-score cell slot for every row so the fallback logic is observable. +const DataTableStub = { + props: ['columns', 'data'], + template: ` +
+
+ +
+
+ ` +} + +function mountView() { + return mount(AccountsView, { + global: { + stubs: { + AppLayout: { template: '
' }, + TablePageLayout: { + template: '
' + }, + DataTable: DataTableStub, + HelpTooltip: true, + Pagination: true, + ConfirmDialog: true, + AccountTableActions: { template: '
' }, + AccountTableFilters: { template: '
' }, + AccountBulkActionsBar: true, + AccountActionMenu: true, + ImportDataModal: true, + ReAuthAccountModal: true, + AccountTestModal: true, + AccountStatsModal: true, + ScheduledTestsPanel: true, + SyncFromCrsModal: true, + TempUnschedStatusModal: true, + ErrorPassthroughRulesModal: true, + TLSFingerprintProfilesModal: true, + CreateAccountModal: true, + EditAccountModal: true, + BulkEditAccountModal: true, + PlatformTypeBadge: true, + AccountCapacityCell: true, + AccountStatusIndicator: true, + AccountTodayStatsCell: true, + AccountGroupsCell: true, + AccountUsageCell: true, + Icon: true + } + } + }) +} + +const baseAccount = { + platform: 'openai', + type: 'apikey', + status: 'active', + schedulable: true, + concurrency: 1, + priority: 0, + error_message: null, + last_used_at: null, + expires_at: null, + auto_pause_on_expired: false, + created_at: '2026-01-01T00:00:00Z', + updated_at: '2026-01-01T00:00:00Z' +} + +describe('admin AccountsView scheduler score column', () => { + beforeEach(() => { + localStorage.clear() + + listAccounts.mockReset() + listWithEtag.mockReset() + getBatchTodayStats.mockReset() + getAllProxies.mockReset() + getAllGroups.mockReset() + + listAccounts.mockResolvedValue({ + items: [ + { + ...baseAccount, + id: 1, + name: 'ungrouped-openai', + // 未分组账号:后端只返回基础分(scheduler_score),无分组维度分数 + scheduler_score: { + base_score: 1.234567, + sticky_score: 0, + sticky_weighted_enabled: false + } + }, + { + ...baseAccount, + id: 2, + name: 'grouped-openai', + scheduler_score: { + base_score: 2, + sticky_score: 3, + sticky_weighted_enabled: true + }, + scheduler_scores: [ + { + group_id: 5, + group_name: 'group-five', + base_score: 2, + sticky_score: 3, + sticky_weighted_enabled: true + } + ] + }, + { + ...baseAccount, + id: 3, + name: 'no-score', + platform: 'anthropic' + } + ], + total: 3, + page: 1, + page_size: 20, + pages: 1 + }) + listWithEtag.mockResolvedValue({ + notModified: true, + etag: null, + data: null + }) + getBatchTodayStats.mockResolvedValue({ stats: {} }) + getAllProxies.mockResolvedValue([]) + getAllGroups.mockResolvedValue([]) + }) + + it('falls back to the base score for ungrouped accounts instead of showing a dash', async () => { + const wrapper = mountView() + await flushPromises() + + const ungroupedCell = wrapper.find('[data-test="scheduler-score-1"]') + expect(ungroupedCell.exists()).toBe(true) + expect(ungroupedCell.text()).toContain('1.234567') + expect(ungroupedCell.text()).toContain('admin.accounts.schedulerScore.ungrouped') + expect(ungroupedCell.text()).not.toBe('-') + }) + + it('renders per-group scores for grouped accounts', async () => { + const wrapper = mountView() + await flushPromises() + + const groupedCell = wrapper.find('[data-test="scheduler-score-2"]') + expect(groupedCell.exists()).toBe(true) + expect(groupedCell.text()).toContain('group-five') + expect(groupedCell.text()).toContain('2') + }) + + it('still shows a dash when no scheduler score is available', async () => { + const wrapper = mountView() + await flushPromises() + + const emptyCell = wrapper.find('[data-test="scheduler-score-3"]') + expect(emptyCell.exists()).toBe(true) + expect(emptyCell.text()).toBe('-') + }) +})