From a8afc209690eb007c745662f488202dc747ab69d Mon Sep 17 00:00:00 2001 From: jjaw Date: Tue, 14 Jul 2026 15:18:47 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=20Ops=20=E9=98=9F=E5=88=97?= =?UTF-8?q?=E6=B7=B1=E5=BA=A6=E8=B4=A6=E5=8F=B7=E6=9F=A5=E8=AF=A2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/internal/repository/account_repo.go | 45 ++++-- ...schedulable_projection_integration_test.go | 107 +++++++++++++++ ...ccount_repo_schedulable_projection_test.go | 82 +++++++++++ .../internal/service/ops_metrics_collector.go | 47 ++++--- .../ops_metrics_collector_projection_test.go | 128 ++++++++++++++++++ 5 files changed, 383 insertions(+), 26 deletions(-) create mode 100644 backend/internal/repository/account_repo_schedulable_projection_integration_test.go create mode 100644 backend/internal/repository/account_repo_schedulable_projection_test.go create mode 100644 backend/internal/service/ops_metrics_collector_projection_test.go diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index 8eb819aeab..d1949299af 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -1050,8 +1050,42 @@ func (r *accountRepository) BindGroups(ctx context.Context, accountID int64, gro } func (r *accountRepository) ListSchedulable(ctx context.Context) ([]service.Account, error) { - now := time.Now() - accounts, err := r.client.Account.Query(). + accounts, err := r.schedulableAccountsQuery(time.Now()).All(ctx) + if err != nil { + return nil, err + } + return r.accountsToService(ctx, accounts) +} + +func (r *accountRepository) ListSchedulableAccountLoads(ctx context.Context) ([]service.AccountWithConcurrency, error) { + accounts, err := r.schedulableAccountsQuery(time.Now()). + Select( + dbaccount.FieldID, + dbaccount.FieldConcurrency, + dbaccount.FieldLoadFactor, + ). + All(ctx) + if err != nil { + return nil, err + } + + loads := make([]service.AccountWithConcurrency, 0, len(accounts)) + for _, account := range accounts { + projection := service.Account{ + ID: account.ID, + Concurrency: account.Concurrency, + LoadFactor: account.LoadFactor, + } + loads = append(loads, service.AccountWithConcurrency{ + ID: account.ID, + MaxConcurrency: projection.EffectiveLoadFactor(), + }) + } + return loads, nil +} + +func (r *accountRepository) schedulableAccountsQuery(now time.Time) *dbent.AccountQuery { + return r.client.Account.Query(). Where( dbaccount.StatusEQ(service.StatusActive), dbaccount.SchedulableEQ(true), @@ -1060,12 +1094,7 @@ func (r *accountRepository) ListSchedulable(ctx context.Context) ([]service.Acco dbaccount.Or(dbaccount.OverloadUntilIsNil(), dbaccount.OverloadUntilLTE(now)), dbaccount.Or(dbaccount.RateLimitResetAtIsNil(), dbaccount.RateLimitResetAtLTE(now)), ). - Order(dbent.Asc(dbaccount.FieldPriority)). - All(ctx) - if err != nil { - return nil, err - } - return r.accountsToService(ctx, accounts) + Order(dbent.Asc(dbaccount.FieldPriority)) } func (r *accountRepository) ListSchedulableByGroupID(ctx context.Context, groupID int64) ([]service.Account, error) { diff --git a/backend/internal/repository/account_repo_schedulable_projection_integration_test.go b/backend/internal/repository/account_repo_schedulable_projection_integration_test.go new file mode 100644 index 0000000000..82d62a1448 --- /dev/null +++ b/backend/internal/repository/account_repo_schedulable_projection_integration_test.go @@ -0,0 +1,107 @@ +//go:build integration + +package repository + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestListSchedulableAccountLoadsMatchesListSchedulable(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + client := tx.Client() + repo := newAccountRepositoryWithSQL(client, tx, nil) + now := time.Now() + past := now.Add(-time.Hour) + future := now.Add(time.Hour) + + create := func(name string) *service.Account { + return mustCreateAccount(t, client, &service.Account{Name: name, Schedulable: true}) + } + + positiveLoad := create("projection-positive-load") + _, err := client.Account.UpdateOneID(positiveLoad.ID).SetConcurrency(2).SetLoadFactor(9).SetPriority(30).Save(ctx) + require.NoError(t, err) + concurrencyFallback := create("projection-concurrency-fallback") + _, err = client.Account.UpdateOneID(concurrencyFallback.ID).SetConcurrency(4).SetPriority(10).Save(ctx) + require.NoError(t, err) + zeroFallback := create("projection-zero-fallback") + _, err = client.Account.UpdateOneID(zeroFallback.ID).SetConcurrency(0).SetLoadFactor(0).SetPriority(20).Save(ctx) + require.NoError(t, err) + + disabled := create("projection-disabled") + _, err = client.Account.UpdateOneID(disabled.ID).SetStatus(service.StatusDisabled).Save(ctx) + require.NoError(t, err) + unschedulable := create("projection-unschedulable") + _, err = client.Account.UpdateOneID(unschedulable.ID).SetSchedulable(false).Save(ctx) + require.NoError(t, err) + expired := create("projection-expired") + _, err = client.Account.UpdateOneID(expired.ID).SetExpiresAt(past).SetAutoPauseOnExpired(true).Save(ctx) + require.NoError(t, err) + expiredAllowed := create("projection-expired-allowed") + _, err = client.Account.UpdateOneID(expiredAllowed.ID).SetExpiresAt(past).SetAutoPauseOnExpired(false).Save(ctx) + require.NoError(t, err) + overloaded := create("projection-overloaded") + _, err = client.Account.UpdateOneID(overloaded.ID).SetOverloadUntil(future).Save(ctx) + require.NoError(t, err) + overloadCleared := create("projection-overload-cleared") + _, err = client.Account.UpdateOneID(overloadCleared.ID).SetOverloadUntil(past).Save(ctx) + require.NoError(t, err) + rateLimited := create("projection-rate-limited") + _, err = client.Account.UpdateOneID(rateLimited.ID).SetRateLimitResetAt(future).Save(ctx) + require.NoError(t, err) + rateLimitCleared := create("projection-rate-limit-cleared") + _, err = client.Account.UpdateOneID(rateLimitCleared.ID).SetRateLimitResetAt(past).Save(ctx) + require.NoError(t, err) + tempBlocked := create("projection-temp-blocked") + _, err = client.Account.UpdateOneID(tempBlocked.ID).SetTempUnschedulableUntil(future).Save(ctx) + require.NoError(t, err) + tempCleared := create("projection-temp-cleared") + _, err = client.Account.UpdateOneID(tempCleared.ID).SetTempUnschedulableUntil(past).Save(ctx) + require.NoError(t, err) + + accounts, err := repo.ListSchedulable(ctx) + require.NoError(t, err) + loads, err := repo.ListSchedulableAccountLoads(ctx) + require.NoError(t, err) + + accountIDs := make([]int64, 0, len(accounts)) + wantByID := make(map[int64]int, len(accounts)) + for i := range accounts { + accountIDs = append(accountIDs, accounts[i].ID) + wantByID[accounts[i].ID] = accounts[i].EffectiveLoadFactor() + } + + loadIDs := make([]int64, 0, len(loads)) + byID := make(map[int64]int, len(loads)) + for _, load := range loads { + loadIDs = append(loadIDs, load.ID) + byID[load.ID] = load.MaxConcurrency + } + require.Equal(t, accountIDs, loadIDs) + targetIDs := map[int64]struct{}{ + positiveLoad.ID: {}, concurrencyFallback.ID: {}, zeroFallback.ID: {}, + } + targetOrder := make([]int64, 0, len(targetIDs)) + for _, id := range loadIDs { + if _, ok := targetIDs[id]; ok { + targetOrder = append(targetOrder, id) + } + } + require.Equal(t, []int64{concurrencyFallback.ID, zeroFallback.ID, positiveLoad.ID}, targetOrder) + require.Equal(t, wantByID, byID) + require.Equal(t, 9, byID[positiveLoad.ID]) + require.Equal(t, 4, byID[concurrencyFallback.ID]) + require.Equal(t, 1, byID[zeroFallback.ID]) + for _, included := range []*service.Account{expiredAllowed, overloadCleared, rateLimitCleared, tempCleared} { + require.Contains(t, byID, included.ID) + } + for _, excluded := range []*service.Account{disabled, unschedulable, expired, overloaded, rateLimited, tempBlocked} { + require.NotContains(t, byID, excluded.ID) + } +} diff --git a/backend/internal/repository/account_repo_schedulable_projection_test.go b/backend/internal/repository/account_repo_schedulable_projection_test.go new file mode 100644 index 0000000000..76d7e0bdc5 --- /dev/null +++ b/backend/internal/repository/account_repo_schedulable_projection_test.go @@ -0,0 +1,82 @@ +package repository + +import ( + "context" + "fmt" + "strings" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + dbent "github.com/Wei-Shaw/sub2api/ent" + _ "github.com/Wei-Shaw/sub2api/ent/runtime" + "github.com/stretchr/testify/require" + + "entgo.io/ent/dialect" + entsql "entgo.io/ent/dialect/sql" +) + +type captureEntQueryMatcher struct { + actual *string +} + +func (m captureEntQueryMatcher) Match(_, actual string) error { + if m.actual == nil { + return fmt.Errorf("query capture target is nil") + } + *m.actual = actual + return nil +} + +func TestListSchedulableAccountLoadsUsesSingleProjectionQuery(t *testing.T) { + var capturedSQL string + db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(captureEntQueryMatcher{actual: &capturedSQL})) + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + + driver := entsql.OpenDB(dialect.Postgres, db) + client := dbent.NewClient(dbent.Driver(driver)) + t.Cleanup(func() { _ = client.Close() }) + repo := newAccountRepositoryWithSQL(client, db, nil) + + mock.ExpectQuery("schedulable account load projection"). + WillReturnRows(sqlmock.NewRows([]string{"id", "concurrency", "load_factor"}). + AddRow(int64(11), 3, nil). + AddRow(int64(12), 2, 7)) + + loads, err := repo.ListSchedulableAccountLoads(context.Background()) + require.NoError(t, err) + require.Len(t, loads, 2) + require.Equal(t, int64(11), loads[0].ID) + require.Equal(t, 3, loads[0].MaxConcurrency) + require.Equal(t, int64(12), loads[1].ID) + require.Equal(t, 7, loads[1].MaxConcurrency) + require.NoError(t, mock.ExpectationsWereMet(), "projection path must execute exactly one query") + + normalized := normalizeSQLWhitespace(capturedSQL) + selectClause, _, found := strings.Cut(normalized, " FROM ") + require.True(t, found, "unexpected projection SQL: %s", normalized) + require.Equal(t, 2, strings.Count(selectClause, ","), "projection must select exactly three columns: %s", selectClause) + require.Contains(t, selectClause, `"id"`) + require.Contains(t, selectClause, `"concurrency"`) + require.Contains(t, selectClause, `"load_factor"`) + require.NotContains(t, selectClause, "credentials") + require.NotContains(t, selectClause, "extra") + require.NotContains(t, selectClause, "proxy_id") + require.NotContains(t, normalized, "account_groups") + require.NotContains(t, normalized, "proxies") + for _, predicateColumn := range []string{ + "status", + "schedulable", + "temp_unschedulable_until", + "expires_at", + "auto_pause_on_expired", + "overload_until", + "rate_limit_reset_at", + "deleted_at", + } { + require.Contains(t, normalized, predicateColumn) + } + _, orderClause, hasOrder := strings.Cut(normalized, " ORDER BY ") + require.True(t, hasOrder, "projection query must preserve schedulable account order: %s", normalized) + require.Contains(t, orderClause, `"priority" ASC`) +} diff --git a/backend/internal/service/ops_metrics_collector.go b/backend/internal/service/ops_metrics_collector.go index b654190cfb..d4f5c32929 100644 --- a/backend/internal/service/ops_metrics_collector.go +++ b/backend/internal/service/ops_metrics_collector.go @@ -39,6 +39,10 @@ const ( var opsMetricsCollectorAdvisoryLockID = hashAdvisoryLockID(opsMetricsCollectorLeaderLockKey) +type opsSchedulableAccountLoadRepository interface { + ListSchedulableAccountLoads(ctx context.Context) ([]AccountWithConcurrency, error) +} + type OpsMetricsCollector struct { opsRepo OpsRepository settingRepo SettingRepository @@ -375,31 +379,16 @@ func (c *OpsMetricsCollector) collectConcurrencyQueueDepth(parentCtx context.Con ctx, cancel := context.WithTimeout(parentCtx, 2*time.Second) defer cancel() - accounts, err := c.accountRepo.ListSchedulable(ctx) + accountLoads, err := c.listSchedulableAccountLoads(ctx) if err != nil { return nil } - if len(accounts) == 0 { + if len(accountLoads) == 0 { zero := 0 return &zero } - batch := make([]AccountWithConcurrency, 0, len(accounts)) - for _, acc := range accounts { - if acc.ID <= 0 { - continue - } - batch = append(batch, AccountWithConcurrency{ - ID: acc.ID, - MaxConcurrency: acc.EffectiveLoadFactor(), - }) - } - if len(batch) == 0 { - zero := 0 - return &zero - } - - loadMap, err := c.concurrencyService.GetAccountsLoadBatch(ctx, batch) + loadMap, err := c.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) if err != nil { return nil } @@ -423,6 +412,28 @@ func (c *OpsMetricsCollector) collectConcurrencyQueueDepth(parentCtx context.Con return &v } +func (c *OpsMetricsCollector) listSchedulableAccountLoads(ctx context.Context) ([]AccountWithConcurrency, error) { + if repo, ok := c.accountRepo.(opsSchedulableAccountLoadRepository); ok { + return repo.ListSchedulableAccountLoads(ctx) + } + + accounts, err := c.accountRepo.ListSchedulable(ctx) + if err != nil { + return nil, err + } + loads := make([]AccountWithConcurrency, 0, len(accounts)) + for _, account := range accounts { + if account.ID <= 0 { + continue + } + loads = append(loads, AccountWithConcurrency{ + ID: account.ID, + MaxConcurrency: account.EffectiveLoadFactor(), + }) + } + return loads, nil +} + type opsCollectedPercentiles struct { p50 *int p90 *int diff --git a/backend/internal/service/ops_metrics_collector_projection_test.go b/backend/internal/service/ops_metrics_collector_projection_test.go new file mode 100644 index 0000000000..e0d1da774d --- /dev/null +++ b/backend/internal/service/ops_metrics_collector_projection_test.go @@ -0,0 +1,128 @@ +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +type opsMetricsProjectionRepo struct { + AccountRepository + accounts []Account + accountLoads []AccountWithConcurrency + listCalls int + projectionCalls int +} + +func (r *opsMetricsProjectionRepo) ListSchedulable(context.Context) ([]Account, error) { + r.listCalls++ + return r.accounts, nil +} + +func (r *opsMetricsProjectionRepo) ListSchedulableAccountLoads(context.Context) ([]AccountWithConcurrency, error) { + r.projectionCalls++ + return r.accountLoads, nil +} + +type opsMetricsFallbackRepo struct { + AccountRepository + accounts []Account + listCalls int +} + +func (r *opsMetricsFallbackRepo) ListSchedulable(context.Context) ([]Account, error) { + r.listCalls++ + return r.accounts, nil +} + +type opsMetricsLoadCache struct { + ConcurrencyCache + loads map[int64]*AccountLoadInfo + got []AccountWithConcurrency +} + +func (c *opsMetricsLoadCache) GetAccountsLoadBatch(_ context.Context, accounts []AccountWithConcurrency) (map[int64]*AccountLoadInfo, error) { + c.got = accounts + return c.loads, nil +} + +func TestCollectConcurrencyQueueDepthUsesProjectionAndPreservesFallbackResult(t *testing.T) { + loadFactor := 7 + accounts := []Account{ + {ID: 11, Concurrency: 2, LoadFactor: &loadFactor}, + {ID: 12, Concurrency: 3}, + {ID: 13}, + } + accountLoads := []AccountWithConcurrency{ + {ID: 11, MaxConcurrency: 7}, + {ID: 12, MaxConcurrency: 3}, + {ID: 13, MaxConcurrency: 1}, + } + loads := map[int64]*AccountLoadInfo{ + 11: {AccountID: 11, WaitingCount: 2}, + 12: {AccountID: 12, WaitingCount: 3}, + 13: {AccountID: 13, WaitingCount: 0}, + } + + projectionRepo := &opsMetricsProjectionRepo{accounts: accounts, accountLoads: accountLoads} + projectionCache := &opsMetricsLoadCache{loads: loads} + projectionConcurrency := NewConcurrencyService(projectionCache) + projectionConcurrency.SetAccountLoadBatchCacheTTL(0) + projectionCollector := &OpsMetricsCollector{ + accountRepo: projectionRepo, + concurrencyService: projectionConcurrency, + } + + fallbackRepo := &opsMetricsFallbackRepo{accounts: accounts} + fallbackCache := &opsMetricsLoadCache{loads: loads} + fallbackConcurrency := NewConcurrencyService(fallbackCache) + fallbackConcurrency.SetAccountLoadBatchCacheTTL(0) + fallbackCollector := &OpsMetricsCollector{ + accountRepo: fallbackRepo, + concurrencyService: fallbackConcurrency, + } + + projectionDepth := projectionCollector.collectConcurrencyQueueDepth(context.Background()) + fallbackDepth := fallbackCollector.collectConcurrencyQueueDepth(context.Background()) + + require.NotNil(t, projectionDepth) + require.NotNil(t, fallbackDepth) + require.Equal(t, 5, *projectionDepth) + require.Equal(t, *fallbackDepth, *projectionDepth) + require.Equal(t, 1, projectionRepo.projectionCalls) + require.Zero(t, projectionRepo.listCalls) + require.Equal(t, 1, fallbackRepo.listCalls) + require.Equal(t, accountLoads, projectionCache.got) + require.Equal(t, fallbackCache.got, projectionCache.got) +} + +func BenchmarkOpsMetricsCollectorCollectConcurrencyQueueDepth(b *testing.B) { + const accountCount = 1000 + loadFactor := 8 + accounts := make([]Account, accountCount) + accountLoads := make([]AccountWithConcurrency, accountCount) + for i := range accountCount { + id := int64(i + 1) + accounts[i] = Account{ + ID: id, + Concurrency: 4, + LoadFactor: &loadFactor, + } + accountLoads[i] = AccountWithConcurrency{ID: id, MaxConcurrency: loadFactor} + } + + repo := &opsMetricsProjectionRepo{accounts: accounts, accountLoads: accountLoads} + cache := &opsMetricsLoadCache{loads: map[int64]*AccountLoadInfo{}} + concurrency := NewConcurrencyService(cache) + concurrency.SetAccountLoadBatchCacheTTL(0) + collector := &OpsMetricsCollector{accountRepo: repo, concurrencyService: concurrency} + + b.ReportAllocs() + b.ResetTimer() + for range b.N { + if depth := collector.collectConcurrencyQueueDepth(context.Background()); depth == nil || *depth != 0 { + b.Fatalf("unexpected queue depth: %v", depth) + } + } +}