mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge pull request #3739 from Wei-Shaw/fix/openai-advanced-scheduler-audit-fixes
fix(scheduler): 修复 OpenAI 高级调度器审计发现的正确性与性能问题
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -174,6 +174,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
service.OpenAIUpstreamTransportHTTPSSE,
|
||||
"",
|
||||
false,
|
||||
false,
|
||||
service.PlatformGrok,
|
||||
)
|
||||
if err != nil {
|
||||
|
||||
@@ -145,6 +145,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
service.OpenAIUpstreamTransportAny,
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
requestPlatform,
|
||||
)
|
||||
if err != nil {
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) 若已存在工具调用上下文则提前返回
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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 == "" {
|
||||
|
||||
@@ -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 = () => {
|
||||
|
||||
@@ -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<typeof import('vue-i18n')>('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: `
|
||||
<div data-test="data-table">
|
||||
<div v-for="row in data" :key="row.id" :data-test="'scheduler-score-' + row.id">
|
||||
<slot name="cell-scheduler_score" :row="row" />
|
||||
</div>
|
||||
</div>
|
||||
`
|
||||
}
|
||||
|
||||
function mountView() {
|
||||
return mount(AccountsView, {
|
||||
global: {
|
||||
stubs: {
|
||||
AppLayout: { template: '<div><slot /></div>' },
|
||||
TablePageLayout: {
|
||||
template: '<div><slot name="filters" /><slot name="table" /><slot name="pagination" /></div>'
|
||||
},
|
||||
DataTable: DataTableStub,
|
||||
HelpTooltip: true,
|
||||
Pagination: true,
|
||||
ConfirmDialog: true,
|
||||
AccountTableActions: { template: '<div><slot name="beforeCreate" /><slot name="after" /></div>' },
|
||||
AccountTableFilters: { template: '<div></div>' },
|
||||
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('-')
|
||||
})
|
||||
})
|
||||
Reference in New Issue
Block a user