Merge pull request #3661 from jianjianai/fix/group-capacity-batch-hotpath

fix: 优化用户组容量汇总批量统计性能
This commit is contained in:
Wesley Liddick
2026-07-04 09:20:02 +08:00
committed by GitHub
4 changed files with 482 additions and 4 deletions
@@ -1061,6 +1061,90 @@ func (r *accountRepository) ListSchedulableByGroupID(ctx context.Context, groupI
})
}
func (r *accountRepository) ListSchedulableCapacityByGroupIDs(ctx context.Context, groupIDs []int64) ([]service.GroupAccountCapacityRow, error) {
groupIDs = uniquePositiveInt64s(groupIDs)
if len(groupIDs) == 0 {
return []service.GroupAccountCapacityRow{}, nil
}
if r.sql == nil {
rows := make([]service.GroupAccountCapacityRow, 0)
for _, groupID := range groupIDs {
accounts, err := r.ListSchedulableByGroupID(ctx, groupID)
if err != nil {
return nil, err
}
for i := range accounts {
acc := &accounts[i]
rows = append(rows, service.GroupAccountCapacityRow{
GroupID: groupID,
AccountID: acc.ID,
Concurrency: acc.Concurrency,
Extra: copyJSONMap(acc.Extra),
SessionWindowStart: acc.SessionWindowStart,
SessionWindowEnd: acc.SessionWindowEnd,
SessionWindowStatus: acc.SessionWindowStatus,
})
}
}
return rows, nil
}
rows, err := r.sql.QueryContext(ctx, `
SELECT
ag.group_id,
a.id AS account_id,
a.concurrency,
COALESCE(a.extra, '{}'::jsonb)::text AS extra,
a.session_window_start,
a.session_window_end,
COALESCE(a.session_window_status, '') AS session_window_status
FROM account_groups ag
JOIN accounts a ON a.id = ag.account_id
WHERE ag.group_id = ANY($1)
AND a.deleted_at IS NULL
AND a.status = $2
AND a.schedulable = TRUE
AND (a.temp_unschedulable_until IS NULL OR a.temp_unschedulable_until <= $3)
AND (a.expires_at IS NULL OR a.expires_at > $3 OR a.auto_pause_on_expired = FALSE)
AND (a.overload_until IS NULL OR a.overload_until <= $3)
AND (a.rate_limit_reset_at IS NULL OR a.rate_limit_reset_at <= $3)
ORDER BY ag.group_id ASC, ag.priority ASC, a.priority ASC, a.id ASC
`, pq.Array(groupIDs), service.StatusActive, time.Now())
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
out := make([]service.GroupAccountCapacityRow, 0)
for rows.Next() {
var row service.GroupAccountCapacityRow
var extraRaw string
if err := rows.Scan(
&row.GroupID,
&row.AccountID,
&row.Concurrency,
&extraRaw,
&row.SessionWindowStart,
&row.SessionWindowEnd,
&row.SessionWindowStatus,
); err != nil {
return nil, err
}
if extraRaw != "" && extraRaw != "null" {
var extra map[string]any
if err := json.Unmarshal([]byte(extraRaw), &extra); err != nil {
return nil, err
}
row.Extra = extra
}
out = append(out, row)
}
if err := rows.Err(); err != nil {
return nil, err
}
return out, nil
}
func (r *accountRepository) ListSchedulableByPlatform(ctx context.Context, platform string) ([]service.Account, error) {
now := time.Now()
accounts, err := r.client.Account.Query().
+43
View File
@@ -466,6 +466,49 @@ func (r *groupRepository) ListActive(ctx context.Context) ([]service.Group, erro
return outGroups, nil
}
func (r *groupRepository) ListActiveIDs(ctx context.Context) ([]int64, error) {
if r.sql != nil {
rows, err := r.sql.QueryContext(ctx, `
SELECT id
FROM groups
WHERE status = $1
AND deleted_at IS NULL
ORDER BY sort_order ASC, id ASC
`, service.StatusActive)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
ids := make([]int64, 0)
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
if err := rows.Err(); err != nil {
return nil, err
}
return ids, nil
}
groups, err := r.client.Group.Query().
Where(group.StatusEQ(service.StatusActive)).
Select(group.FieldID).
Order(dbent.Asc(group.FieldSortOrder), dbent.Asc(group.FieldID)).
All(ctx)
if err != nil {
return nil, err
}
ids := make([]int64, 0, len(groups))
for i := range groups {
ids = append(ids, groups[i].ID)
}
return ids, nil
}
func (r *groupRepository) ListActiveByPlatform(ctx context.Context, platform string) ([]service.Group, error) {
groups, err := r.client.Group.Query().
Where(group.StatusEQ(service.StatusActive), group.PlatformEQ(platform)).
@@ -16,6 +16,26 @@ type GroupCapacitySummary struct {
RPMMax int `json:"rpm_max"`
}
// GroupAccountCapacityRow is the lightweight account projection needed for
// capacity summary aggregation.
type GroupAccountCapacityRow struct {
GroupID int64
AccountID int64
Concurrency int
Extra map[string]any
SessionWindowStart *time.Time
SessionWindowEnd *time.Time
SessionWindowStatus string
}
type groupCapacityActiveGroupIDLister interface {
ListActiveIDs(ctx context.Context) ([]int64, error)
}
type groupCapacityAccountLister interface {
ListSchedulableCapacityByGroupIDs(ctx context.Context, groupIDs []int64) ([]GroupAccountCapacityRow, error)
}
// GroupCapacityService aggregates per-group capacity from runtime data.
type GroupCapacityService struct {
accountRepo AccountRepository
@@ -44,24 +64,176 @@ func NewGroupCapacityService(
// GetAllGroupCapacity returns capacity summary for all active groups.
func (s *GroupCapacityService) GetAllGroupCapacity(ctx context.Context) ([]GroupCapacitySummary, error) {
groups, err := s.groupRepo.ListActive(ctx)
groupIDs, err := s.listActiveGroupIDs(ctx)
if err != nil {
return nil, err
}
results := make([]GroupCapacitySummary, 0, len(groups))
if lister, ok := s.accountRepo.(groupCapacityAccountLister); ok {
return s.getGroupCapacitiesBatch(ctx, groupIDs, lister)
}
return s.getGroupCapacitiesSequential(ctx, groupIDs), nil
}
func (s *GroupCapacityService) listActiveGroupIDs(ctx context.Context) ([]int64, error) {
if lister, ok := s.groupRepo.(groupCapacityActiveGroupIDLister); ok {
return lister.ListActiveIDs(ctx)
}
groups, err := s.groupRepo.ListActive(ctx)
if err != nil {
return nil, err
}
groupIDs := make([]int64, 0, len(groups))
for i := range groups {
cap, err := s.getGroupCapacity(ctx, groups[i].ID)
groupIDs = append(groupIDs, groups[i].ID)
}
return groupIDs, nil
}
func (s *GroupCapacityService) getGroupCapacitiesSequential(ctx context.Context, groupIDs []int64) []GroupCapacitySummary {
results := make([]GroupCapacitySummary, 0, len(groupIDs))
for _, groupID := range groupIDs {
cap, err := s.getGroupCapacity(ctx, groupID)
if err != nil {
// Skip groups with errors, return partial results
continue
}
cap.GroupID = groups[i].ID
cap.GroupID = groupID
results = append(results, cap)
}
return results
}
type groupCapacityAccountRef struct {
groupID int64
accountID int64
}
func (s *GroupCapacityService) getGroupCapacitiesBatch(ctx context.Context, groupIDs []int64, lister groupCapacityAccountLister) ([]GroupCapacitySummary, error) {
results := make([]GroupCapacitySummary, len(groupIDs))
groupIndex := make(map[int64]int, len(groupIDs))
for i, groupID := range groupIDs {
results[i].GroupID = groupID
groupIndex[groupID] = i
}
if len(groupIDs) == 0 {
return results, nil
}
rows, err := lister.ListSchedulableCapacityByGroupIDs(ctx, groupIDs)
if err != nil {
return nil, err
}
if len(rows) == 0 {
return results, nil
}
refs := make([]groupCapacityAccountRef, 0, len(rows))
seenGroupAccount := make(map[groupCapacityAccountRef]struct{}, len(rows))
accountIDSet := make(map[int64]struct{}, len(rows))
accountIDs := make([]int64, 0, len(rows))
sessionTimeouts := make(map[int64]time.Duration)
for _, row := range rows {
idx, ok := groupIndex[row.GroupID]
if !ok || row.AccountID <= 0 {
continue
}
ref := groupCapacityAccountRef{groupID: row.GroupID, accountID: row.AccountID}
if _, ok := seenGroupAccount[ref]; ok {
continue
}
seenGroupAccount[ref] = struct{}{}
refs = append(refs, ref)
if _, ok := accountIDSet[row.AccountID]; !ok {
accountIDSet[row.AccountID] = struct{}{}
accountIDs = append(accountIDs, row.AccountID)
}
acc := Account{
ID: row.AccountID,
Concurrency: row.Concurrency,
Extra: row.Extra,
SessionWindowStart: row.SessionWindowStart,
SessionWindowEnd: row.SessionWindowEnd,
SessionWindowStatus: row.SessionWindowStatus,
}
results[idx].ConcurrencyMax += acc.Concurrency
if maxSessions := acc.GetMaxSessions(); maxSessions > 0 {
results[idx].SessionsMax += maxSessions
timeout := time.Duration(acc.GetSessionIdleTimeoutMinutes()) * time.Minute
if timeout <= 0 {
timeout = 5 * time.Minute
}
sessionTimeouts[acc.ID] = timeout
}
if rpm := acc.GetBaseRPM(); rpm > 0 {
results[idx].RPMMax += rpm
}
}
if len(accountIDs) == 0 {
return results, nil
}
concurrencyMap := map[int64]int{}
if s.concurrencyService != nil {
concurrencyMap, _ = s.concurrencyService.GetAccountConcurrencyBatch(ctx, accountIDs)
}
sessionAccountIDs := accountIDsForGroupsWithLimit(refs, groupIndex, results, func(summary GroupCapacitySummary) bool {
return summary.SessionsMax > 0
})
var sessionsMap map[int64]int
if len(sessionAccountIDs) > 0 && s.sessionLimitCache != nil {
sessionsMap, _ = s.sessionLimitCache.GetActiveSessionCountBatch(ctx, sessionAccountIDs, sessionTimeouts)
}
rpmAccountIDs := accountIDsForGroupsWithLimit(refs, groupIndex, results, func(summary GroupCapacitySummary) bool {
return summary.RPMMax > 0
})
var rpmMap map[int64]int
if len(rpmAccountIDs) > 0 && s.rpmCache != nil {
rpmMap, _ = s.rpmCache.GetRPMBatch(ctx, rpmAccountIDs)
}
for _, ref := range refs {
idx := groupIndex[ref.groupID]
results[idx].ConcurrencyUsed += concurrencyMap[ref.accountID]
if sessionsMap != nil && results[idx].SessionsMax > 0 {
results[idx].SessionsUsed += sessionsMap[ref.accountID]
}
if rpmMap != nil && results[idx].RPMMax > 0 {
results[idx].RPMUsed += rpmMap[ref.accountID]
}
}
return results, nil
}
func accountIDsForGroupsWithLimit(refs []groupCapacityAccountRef, groupIndex map[int64]int, summaries []GroupCapacitySummary, include func(GroupCapacitySummary) bool) []int64 {
seen := make(map[int64]struct{})
accountIDs := make([]int64, 0)
for _, ref := range refs {
idx, ok := groupIndex[ref.groupID]
if !ok || !include(summaries[idx]) {
continue
}
if _, ok := seen[ref.accountID]; ok {
continue
}
seen[ref.accountID] = struct{}{}
accountIDs = append(accountIDs, ref.accountID)
}
return accountIDs
}
func (s *GroupCapacityService) getGroupCapacity(ctx context.Context, groupID int64) (GroupCapacitySummary, error) {
accounts, err := s.accountRepo.ListSchedulableByGroupID(ctx, groupID)
if err != nil {
@@ -0,0 +1,179 @@
package service
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/require"
)
type groupCapacityAccountRepoStub struct {
AccountRepository
rows []GroupAccountCapacityRow
requested []int64
}
func (s *groupCapacityAccountRepoStub) ListSchedulableCapacityByGroupIDs(_ context.Context, groupIDs []int64) ([]GroupAccountCapacityRow, error) {
s.requested = append([]int64(nil), groupIDs...)
return append([]GroupAccountCapacityRow(nil), s.rows...), nil
}
type groupCapacityGroupRepoStub struct {
GroupRepository
groupIDs []int64
listCalls int
}
func (s *groupCapacityGroupRepoStub) ListActiveIDs(context.Context) ([]int64, error) {
s.listCalls++
return append([]int64(nil), s.groupIDs...), nil
}
type groupCapacityConcurrencyCacheStub struct {
ConcurrencyCache
counts map[int64]int
requested []int64
}
func (s *groupCapacityConcurrencyCacheStub) GetAccountConcurrencyBatch(_ context.Context, accountIDs []int64) (map[int64]int, error) {
s.requested = append([]int64(nil), accountIDs...)
out := make(map[int64]int, len(accountIDs))
for _, id := range accountIDs {
out[id] = s.counts[id]
}
return out, nil
}
type groupCapacitySessionCacheStub struct {
SessionLimitCache
counts map[int64]int
requested []int64
idleTimeouts map[int64]time.Duration
}
func (s *groupCapacitySessionCacheStub) GetActiveSessionCountBatch(_ context.Context, accountIDs []int64, idleTimeouts map[int64]time.Duration) (map[int64]int, error) {
s.requested = append([]int64(nil), accountIDs...)
s.idleTimeouts = make(map[int64]time.Duration, len(idleTimeouts))
for id, timeout := range idleTimeouts {
s.idleTimeouts[id] = timeout
}
out := make(map[int64]int, len(accountIDs))
for _, id := range accountIDs {
out[id] = s.counts[id]
}
return out, nil
}
type groupCapacityRPMCacheStub struct {
RPMCache
counts map[int64]int
requested []int64
}
func (s *groupCapacityRPMCacheStub) GetRPMBatch(_ context.Context, accountIDs []int64) (map[int64]int, error) {
s.requested = append([]int64(nil), accountIDs...)
out := make(map[int64]int, len(accountIDs))
for _, id := range accountIDs {
out[id] = s.counts[id]
}
return out, nil
}
func TestGetAllGroupCapacityBatchAggregatesRuntimeAndLimits(t *testing.T) {
accountRepo := &groupCapacityAccountRepoStub{
rows: []GroupAccountCapacityRow{
{
GroupID: 10,
AccountID: 1,
Concurrency: 2,
Extra: map[string]any{
"max_sessions": 3,
"session_idle_timeout_minutes": 7,
"base_rpm": 11,
},
},
{
GroupID: 20,
AccountID: 1,
Concurrency: 2,
Extra: map[string]any{
"max_sessions": 3,
"session_idle_timeout_minutes": 7,
"base_rpm": 11,
},
},
{
GroupID: 20,
AccountID: 2,
Concurrency: 4,
Extra: map[string]any{
"max_sessions": 1,
"session_idle_timeout_minutes": 9,
"base_rpm": 13,
},
},
},
}
groupRepo := &groupCapacityGroupRepoStub{groupIDs: []int64{10, 20}}
concurrencyCache := &groupCapacityConcurrencyCacheStub{counts: map[int64]int{1: 1, 2: 2}}
sessionCache := &groupCapacitySessionCacheStub{counts: map[int64]int{1: 2, 2: 1}}
rpmCache := &groupCapacityRPMCacheStub{counts: map[int64]int{1: 5, 2: 7}}
svc := NewGroupCapacityService(
accountRepo,
groupRepo,
NewConcurrencyService(concurrencyCache),
sessionCache,
rpmCache,
)
results, err := svc.GetAllGroupCapacity(context.Background())
require.NoError(t, err)
require.Equal(t, 1, groupRepo.listCalls)
require.Equal(t, []int64{10, 20}, accountRepo.requested)
require.Equal(t, []int64{1, 2}, concurrencyCache.requested)
require.ElementsMatch(t, []int64{1, 2}, sessionCache.requested)
require.ElementsMatch(t, []int64{1, 2}, rpmCache.requested)
require.Equal(t, 7*time.Minute, sessionCache.idleTimeouts[1])
require.Equal(t, 9*time.Minute, sessionCache.idleTimeouts[2])
require.Equal(t, []GroupCapacitySummary{
{
GroupID: 10,
ConcurrencyUsed: 1,
ConcurrencyMax: 2,
SessionsUsed: 2,
SessionsMax: 3,
RPMUsed: 5,
RPMMax: 11,
},
{
GroupID: 20,
ConcurrencyUsed: 3,
ConcurrencyMax: 6,
SessionsUsed: 3,
SessionsMax: 4,
RPMUsed: 12,
RPMMax: 24,
},
}, results)
}
func TestGetAllGroupCapacityBatchKeepsEmptyGroupRows(t *testing.T) {
accountRepo := &groupCapacityAccountRepoStub{
rows: []GroupAccountCapacityRow{
{GroupID: 20, AccountID: 2, Concurrency: 4},
},
}
groupRepo := &groupCapacityGroupRepoStub{groupIDs: []int64{10, 20}}
svc := NewGroupCapacityService(accountRepo, groupRepo, nil, nil, nil)
results, err := svc.GetAllGroupCapacity(context.Background())
require.NoError(t, err)
require.Equal(t, []GroupCapacitySummary{
{GroupID: 10},
{GroupID: 20, ConcurrencyMax: 4},
}, results)
}