diff --git a/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go b/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go index 75e3c1aa8c..685231fc91 100644 --- a/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go +++ b/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go @@ -31,9 +31,18 @@ type fakeSchedulerCache struct { func (f *fakeSchedulerCache) GetSnapshot(_ context.Context, _ service.SchedulerBucket) ([]*service.Account, bool, error) { return f.accounts, true, nil } -func (f *fakeSchedulerCache) SetSnapshot(_ context.Context, _ service.SchedulerBucket, _ []service.Account) error { +func (f *fakeSchedulerCache) CaptureBucketWriteToken(_ context.Context, bucket service.SchedulerBucket) (service.SchedulerBucketWriteToken, error) { + return service.SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} +func (f *fakeSchedulerCache) SetSnapshot(_ context.Context, _ service.SchedulerBucket, _ service.SchedulerBucketWriteToken, _ []service.Account) error { return nil } +func (f *fakeSchedulerCache) RetireBucket(_ context.Context, _ service.SchedulerBucket) error { + return nil +} +func (f *fakeSchedulerCache) ReopenBucket(_ context.Context, bucket service.SchedulerBucket) (service.SchedulerBucketWriteToken, error) { + return service.SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} func (f *fakeSchedulerCache) GetAccount(_ context.Context, id int64) (*service.Account, error) { for _, account := range f.accounts { if account != nil && account.ID == id { diff --git a/backend/internal/repository/account_repo_integration_test.go b/backend/internal/repository/account_repo_integration_test.go index f82ad4ca28..7c928072e4 100644 --- a/backend/internal/repository/account_repo_integration_test.go +++ b/backend/internal/repository/account_repo_integration_test.go @@ -31,10 +31,22 @@ func (s *schedulerCacheRecorder) GetSnapshot(ctx context.Context, bucket service return nil, false, nil } -func (s *schedulerCacheRecorder) SetSnapshot(ctx context.Context, bucket service.SchedulerBucket, accounts []service.Account) error { +func (s *schedulerCacheRecorder) CaptureBucketWriteToken(ctx context.Context, bucket service.SchedulerBucket) (service.SchedulerBucketWriteToken, error) { + return service.SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} + +func (s *schedulerCacheRecorder) SetSnapshot(ctx context.Context, bucket service.SchedulerBucket, token service.SchedulerBucketWriteToken, accounts []service.Account) error { return nil } +func (s *schedulerCacheRecorder) RetireBucket(ctx context.Context, bucket service.SchedulerBucket) error { + return nil +} + +func (s *schedulerCacheRecorder) ReopenBucket(ctx context.Context, bucket service.SchedulerBucket) (service.SchedulerBucketWriteToken, error) { + return service.SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} + func (s *schedulerCacheRecorder) GetAccount(ctx context.Context, accountID int64) (*service.Account, error) { if s.accounts == nil { return nil, nil diff --git a/backend/internal/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go index c68bcf96e6..72d12d799b 100644 --- a/backend/internal/repository/scheduler_cache.go +++ b/backend/internal/repository/scheduler_cache.go @@ -20,6 +20,8 @@ const ( schedulerActivePrefix = "sched:active:" schedulerReadyPrefix = "sched:ready:" schedulerVersionPrefix = "sched:ver:" + schedulerEpochPrefix = "sched:epoch:" + schedulerRetiredPrefix = "sched:retired:" schedulerSnapshotPrefix = "sched:" schedulerLockPrefix = "sched:lock:" @@ -32,6 +34,98 @@ const ( ) var ( + captureBucketWriteTokenScript = redis.NewScript(` +if redis.call('EXISTS', KEYS[2]) == 1 then + return -1 +end + +local currentEpoch = redis.call('GET', KEYS[1]) +if currentEpoch == false then + redis.call('SET', KEYS[1], '1') + return 1 +end + +local parsedEpoch = tonumber(currentEpoch) +if parsedEpoch == nil or parsedEpoch < 1 then + return -2 +end +return parsedEpoch +`) + + allocateSnapshotVersionScript = redis.NewScript(` +if redis.call('EXISTS', KEYS[2]) == 1 then + return -1 +end + +local currentEpoch = tonumber(redis.call('GET', KEYS[1])) +local expectedEpoch = tonumber(ARGV[1]) +if currentEpoch == nil or expectedEpoch == nil or currentEpoch ~= expectedEpoch then + return -2 +end + +return redis.call('INCR', KEYS[3]) +`) + + retireBucketScript = redis.NewScript(` +local retired = redis.call('GET', KEYS[2]) +local currentEpoch = tonumber(redis.call('GET', KEYS[1])) or 0 + +if retired == false then + currentEpoch = currentEpoch + 1 + if currentEpoch < 1 then + currentEpoch = 1 + end + redis.call('SET', KEYS[1], tostring(currentEpoch)) + redis.call('SET', KEYS[2], tostring(currentEpoch)) +elseif currentEpoch < 1 then + currentEpoch = tonumber(retired) or 1 + redis.call('SET', KEYS[1], tostring(currentEpoch)) +end + +redis.call('SREM', KEYS[3], ARGV[1]) +local currentActive = redis.call('GET', KEYS[5]) +if currentActive ~= false then + redis.call('EXPIRE', ARGV[2] .. currentActive, tonumber(ARGV[3])) +end +redis.call('DEL', KEYS[4], KEYS[5]) +return currentEpoch +`) + + reopenBucketScript = redis.NewScript(` +local currentEpochRaw = redis.call('GET', KEYS[1]) +local currentEpoch = tonumber(currentEpochRaw) +local retiredEpochRaw = redis.call('GET', KEYS[2]) + +if retiredEpochRaw == false then + if currentEpochRaw == false then + redis.call('SET', KEYS[1], '1') + return 1 + end + if currentEpoch == nil or currentEpoch < 1 then + return -2 + end + return currentEpoch +end + +local retiredEpoch = tonumber(retiredEpochRaw) +if retiredEpoch == nil or retiredEpoch < 1 then + return -2 +end +if currentEpoch == nil or currentEpoch < retiredEpoch then + currentEpoch = retiredEpoch +end + +redis.call('SET', KEYS[1], tostring(currentEpoch)) +redis.call('DEL', KEYS[2]) +redis.call('SREM', KEYS[3], ARGV[1]) +local currentActive = redis.call('GET', KEYS[5]) +if currentActive ~= false then + redis.call('EXPIRE', ARGV[2] .. currentActive, tonumber(ARGV[3])) +end +redis.call('DEL', KEYS[4], KEYS[5]) +return currentEpoch +`) + // activateSnapshotScript 原子 CAS 切换快照版本。 // 仅当新版本号 >= 当前激活版本时才切换,防止并发写入导致版本回滚。 // 旧快照使用 EXPIRE 设置宽限期而非立即 DEL,避免与 reader 竞态。 @@ -40,13 +134,28 @@ var ( // KEYS[2] = readyKey (sched:ready:{bucket}) // KEYS[3] = bucketSetKey (sched:buckets) // KEYS[4] = snapshotKey (新写入的快照 key) + // KEYS[5] = epochKey + // KEYS[6] = retiredKey // ARGV[1] = 新版本号字符串 // ARGV[2] = bucket 字符串 (用于 SADD) // ARGV[3] = 快照 key 前缀 (用于构造旧快照 key) // ARGV[4] = 宽限期 TTL 秒数 + // ARGV[5] = writer epoch // // 返回 1 = 已激活, 0 = 版本过旧未激活 activateSnapshotScript = redis.NewScript(` +if redis.call('EXISTS', KEYS[6]) == 1 then + redis.call('DEL', KEYS[4]) + return -1 +end + +local currentEpoch = tonumber(redis.call('GET', KEYS[5])) +local expectedEpoch = tonumber(ARGV[5]) +if currentEpoch == nil or expectedEpoch == nil or currentEpoch ~= expectedEpoch then + redis.call('DEL', KEYS[4]) + return -2 +end + local currentActive = redis.call('GET', KEYS[1]) local newVersion = tonumber(ARGV[1]) @@ -151,19 +260,87 @@ func (c *schedulerCache) GetSnapshot(ctx context.Context, bucket service.Schedul return accounts, true, nil } -func (c *schedulerCache) SetSnapshot(ctx context.Context, bucket service.SchedulerBucket, accounts []service.Account) error { - // Phase 1: 分配新版本号并写入快照数据。 - // INCR 保证每个调用方获得唯一递增版本号。 - // 写入的 snapshotKey 是新的版本化 key,reader 尚不知晓,因此无竞态。 - versionKey := schedulerBucketKey(schedulerVersionPrefix, bucket) - version, err := c.rdb.Incr(ctx, versionKey).Result() +func (c *schedulerCache) CaptureBucketWriteToken(ctx context.Context, bucket service.SchedulerBucket) (service.SchedulerBucketWriteToken, error) { + result, err := captureBucketWriteTokenScript.Run(ctx, c.rdb, []string{ + schedulerBucketKey(schedulerEpochPrefix, bucket), + schedulerBucketKey(schedulerRetiredPrefix, bucket), + }).Int64() + if err != nil { + return service.SchedulerBucketWriteToken{}, err + } + if err := schedulerBucketWriteResultError(result, bucket); err != nil { + return service.SchedulerBucketWriteToken{}, err + } + return service.SchedulerBucketWriteToken{Bucket: bucket, Epoch: result}, nil +} + +func (c *schedulerCache) RetireBucket(ctx context.Context, bucket service.SchedulerBucket) error { + snapshotKeyPrefix := fmt.Sprintf("%s%d:%s:%s:v", schedulerSnapshotPrefix, bucket.GroupID, bucket.Platform, bucket.Mode) + result, err := retireBucketScript.Run(ctx, c.rdb, []string{ + schedulerBucketKey(schedulerEpochPrefix, bucket), + schedulerBucketKey(schedulerRetiredPrefix, bucket), + schedulerBucketSetKey, + schedulerBucketKey(schedulerReadyPrefix, bucket), + schedulerBucketKey(schedulerActivePrefix, bucket), + }, bucket.String(), snapshotKeyPrefix, snapshotGraceTTLSeconds).Int64() if err != nil { return err } + if result < 1 { + return fmt.Errorf("retire scheduler bucket %s returned invalid epoch %d", bucket.String(), result) + } + return nil +} - versionStr := strconv.FormatInt(version, 10) - snapshotKey := schedulerSnapshotKey(bucket, versionStr) +func (c *schedulerCache) ReopenBucket(ctx context.Context, bucket service.SchedulerBucket) (service.SchedulerBucketWriteToken, error) { + snapshotKeyPrefix := fmt.Sprintf("%s%d:%s:%s:v", schedulerSnapshotPrefix, bucket.GroupID, bucket.Platform, bucket.Mode) + result, err := reopenBucketScript.Run(ctx, c.rdb, []string{ + schedulerBucketKey(schedulerEpochPrefix, bucket), + schedulerBucketKey(schedulerRetiredPrefix, bucket), + schedulerBucketSetKey, + schedulerBucketKey(schedulerReadyPrefix, bucket), + schedulerBucketKey(schedulerActivePrefix, bucket), + }, bucket.String(), snapshotKeyPrefix, snapshotGraceTTLSeconds).Int64() + if err != nil { + return service.SchedulerBucketWriteToken{}, err + } + if err := schedulerBucketWriteResultError(result, bucket); err != nil { + return service.SchedulerBucketWriteToken{}, err + } + return service.SchedulerBucketWriteToken{Bucket: bucket, Epoch: result}, nil +} +func (c *schedulerCache) SetSnapshot(ctx context.Context, bucket service.SchedulerBucket, token service.SchedulerBucketWriteToken, accounts []service.Account) error { + if !token.ValidFor(bucket) { + return fmt.Errorf("%w: bucket=%s", service.ErrSchedulerBucketWriteFenced, bucket.String()) + } + version, err := c.allocateSnapshotVersion(ctx, bucket, token) + if err != nil { + return err + } + if err := c.writeSnapshotVersion(ctx, bucket, version, accounts); err != nil { + return err + } + return c.activateSnapshotVersion(ctx, bucket, token, version) +} + +func (c *schedulerCache) allocateSnapshotVersion(ctx context.Context, bucket service.SchedulerBucket, token service.SchedulerBucketWriteToken) (string, error) { + result, err := allocateSnapshotVersionScript.Run(ctx, c.rdb, []string{ + schedulerBucketKey(schedulerEpochPrefix, bucket), + schedulerBucketKey(schedulerRetiredPrefix, bucket), + schedulerBucketKey(schedulerVersionPrefix, bucket), + }, token.Epoch).Int64() + if err != nil { + return "", err + } + if err := schedulerBucketWriteResultError(result, bucket); err != nil { + return "", err + } + return strconv.FormatInt(result, 10), nil +} + +func (c *schedulerCache) writeSnapshotVersion(ctx context.Context, bucket service.SchedulerBucket, version string, accounts []service.Account) error { + snapshotKey := schedulerSnapshotKey(bucket, version) cacheableAccounts, err := c.writeAccounts(ctx, accounts) if err != nil { return err @@ -191,7 +368,12 @@ func (c *schedulerCache) SetSnapshot(ctx context.Context, bucket service.Schedul } } - // Phase 2: 原子 CAS 激活版本。 + return nil +} + +func (c *schedulerCache) activateSnapshotVersion(ctx context.Context, bucket service.SchedulerBucket, token service.SchedulerBucketWriteToken, version string) error { + snapshotKey := schedulerSnapshotKey(bucket, version) + // Phase 2: 原子 CAS 切换版本,同时再次校验退休状态与 writer epoch。 // Lua 脚本保证:仅当新版本 >= 当前激活版本时才切换 active 指针, // 防止并发写入导致版本回滚。 // 旧快照使用 EXPIRE 宽限期而非立即 DEL,避免 reader 竞态。 @@ -199,15 +381,32 @@ func (c *schedulerCache) SetSnapshot(ctx context.Context, bucket service.Schedul readyKey := schedulerBucketKey(schedulerReadyPrefix, bucket) snapshotKeyPrefix := fmt.Sprintf("%s%d:%s:%s:v", schedulerSnapshotPrefix, bucket.GroupID, bucket.Platform, bucket.Mode) - keys := []string{activeKey, readyKey, schedulerBucketSetKey, snapshotKey} - args := []any{versionStr, bucket.String(), snapshotKeyPrefix, snapshotGraceTTLSeconds} + keys := []string{ + activeKey, + readyKey, + schedulerBucketSetKey, + snapshotKey, + schedulerBucketKey(schedulerEpochPrefix, bucket), + schedulerBucketKey(schedulerRetiredPrefix, bucket), + } + args := []any{version, bucket.String(), snapshotKeyPrefix, snapshotGraceTTLSeconds, token.Epoch} - _, err = activateSnapshotScript.Run(ctx, c.rdb, keys, args...).Result() + result, err := activateSnapshotScript.Run(ctx, c.rdb, keys, args...).Int64() if err != nil { return err } + return schedulerBucketWriteResultError(result, bucket) +} - return nil +func schedulerBucketWriteResultError(result int64, bucket service.SchedulerBucket) error { + switch result { + case -1: + return fmt.Errorf("%w: bucket=%s", service.ErrSchedulerBucketRetired, bucket.String()) + case -2: + return fmt.Errorf("%w: bucket=%s", service.ErrSchedulerBucketWriteFenced, bucket.String()) + default: + return nil + } } func (c *schedulerCache) GetAccount(ctx context.Context, accountID int64) (*service.Account, error) { diff --git a/backend/internal/repository/scheduler_cache_integration_test.go b/backend/internal/repository/scheduler_cache_integration_test.go index 948c2c73e0..00966d95b6 100644 --- a/backend/internal/repository/scheduler_cache_integration_test.go +++ b/backend/internal/repository/scheduler_cache_integration_test.go @@ -67,7 +67,9 @@ func TestSchedulerCacheSnapshotUsesSlimMetadataButKeepsFullAccount(t *testing.T) }, } - require.NoError(t, cache.SetSnapshot(ctx, bucket, []service.Account{account})) + token, err := cache.CaptureBucketWriteToken(ctx, bucket) + require.NoError(t, err) + require.NoError(t, cache.SetSnapshot(ctx, bucket, token, []service.Account{account})) snapshot, hit, err := cache.GetSnapshot(ctx, bucket) require.NoError(t, err) @@ -102,3 +104,36 @@ func TestSchedulerCacheSnapshotUsesSlimMetadataButKeepsFullAccount(t *testing.T) require.Len(t, full.AccountGroups, 1) require.NotNil(t, full.AccountGroups[0].Group) } + +func TestSchedulerCacheRetireAndReopenFencesOldEpochIntegration(t *testing.T) { + ctx := context.Background() + rdb := testRedis(t) + cache := NewSchedulerCache(rdb) + bucket := service.SchedulerBucket{GroupID: 77, Platform: service.PlatformAntigravity, Mode: service.SchedulerModeForced} + account := service.Account{ID: 7701, Platform: service.PlatformAntigravity, Type: service.AccountTypeOAuth} + + oldToken, err := cache.CaptureBucketWriteToken(ctx, bucket) + require.NoError(t, err) + require.NoError(t, cache.SetSnapshot(ctx, bucket, oldToken, []service.Account{account})) + require.NoError(t, cache.RetireBucket(ctx, bucket)) + require.NoError(t, cache.RetireBucket(ctx, bucket)) + + _, hit, err := cache.GetSnapshot(ctx, bucket) + require.NoError(t, err) + require.False(t, hit) + _, err = cache.CaptureBucketWriteToken(ctx, bucket) + require.ErrorIs(t, err, service.ErrSchedulerBucketRetired) + require.ErrorIs(t, cache.SetSnapshot(ctx, bucket, oldToken, []service.Account{account}), service.ErrSchedulerBucketRetired) + + newToken, err := cache.ReopenBucket(ctx, bucket) + require.NoError(t, err) + require.Greater(t, newToken.Epoch, oldToken.Epoch) + require.ErrorIs(t, cache.SetSnapshot(ctx, bucket, oldToken, []service.Account{account}), service.ErrSchedulerBucketWriteFenced) + require.NoError(t, cache.SetSnapshot(ctx, bucket, newToken, []service.Account{account})) + + snapshot, hit, err := cache.GetSnapshot(ctx, bucket) + require.NoError(t, err) + require.True(t, hit) + require.Len(t, snapshot, 1) + require.Equal(t, account.ID, snapshot[0].ID) +} diff --git a/backend/internal/repository/scheduler_cache_unit_test.go b/backend/internal/repository/scheduler_cache_unit_test.go index ecca7f3892..bcba575145 100644 --- a/backend/internal/repository/scheduler_cache_unit_test.go +++ b/backend/internal/repository/scheduler_cache_unit_test.go @@ -14,13 +14,18 @@ import ( ) func newSchedulerCacheUnit(t *testing.T) *schedulerCache { + cache, _ := newSchedulerCacheUnitWithRedis(t) + return cache +} + +func newSchedulerCacheUnitWithRedis(t *testing.T) (*schedulerCache, *miniredis.Miniredis) { t.Helper() mr := miniredis.RunT(t) rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) t.Cleanup(func() { _ = rdb.Close() }) cache, ok := newSchedulerCacheWithChunkSizes(rdb, defaultSchedulerSnapshotMGetChunkSize, defaultSchedulerSnapshotWriteChunkSize).(*schedulerCache) require.True(t, ok) - return cache + return cache, mr } func TestSchedulerCacheWriteAccountsSkipsUnencodableTimes(t *testing.T) { @@ -229,3 +234,207 @@ func TestBuildSchedulerMetadataAccount_KeepsSparkShadowRoutingIdentity(t *testin require.Equal(t, map[string]any{"gpt-5.4": "gpt-5.4-openai-compact"}, got.Credentials["compact_model_mapping"]) require.Nil(t, got.Credentials["access_token"]) } + +func TestSchedulerCacheBucketRetirementFencesWritersAndReopen(t *testing.T) { + ctx := context.Background() + cache, mr := newSchedulerCacheUnitWithRedis(t) + bucket := service.SchedulerBucket{GroupID: 41, Platform: service.PlatformOpenAI, Mode: service.SchedulerModeSingle} + otherBucket := service.SchedulerBucket{GroupID: 42, Platform: service.PlatformOpenAI, Mode: service.SchedulerModeSingle} + account := service.Account{ID: 4101, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey} + + token, err := cache.CaptureBucketWriteToken(ctx, bucket) + require.NoError(t, err) + require.True(t, token.ValidFor(bucket)) + require.NoError(t, cache.SetSnapshot(ctx, bucket, token, []service.Account{account})) + + // A token is bound to the full bucket identity, not just an epoch number. + err = cache.SetSnapshot(ctx, otherBucket, token, []service.Account{account}) + require.ErrorIs(t, err, service.ErrSchedulerBucketWriteFenced) + _, err = cache.rdb.Get(ctx, schedulerBucketKey(schedulerVersionPrefix, otherBucket)).Result() + require.ErrorIs(t, err, redis.Nil) + otherAccount := service.Account{ID: 4201, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey} + otherToken, err := cache.CaptureBucketWriteToken(ctx, otherBucket) + require.NoError(t, err) + require.NoError(t, cache.SetSnapshot(ctx, otherBucket, otherToken, []service.Account{otherAccount})) + otherEpoch := otherToken.Epoch + + activeVersion, err := cache.rdb.Get(ctx, schedulerBucketKey(schedulerActivePrefix, bucket)).Result() + require.NoError(t, err) + require.NoError(t, cache.RetireBucket(ctx, bucket)) + retiredEpoch, err := cache.rdb.Get(ctx, schedulerBucketKey(schedulerEpochPrefix, bucket)).Int64() + require.NoError(t, err) + require.Greater(t, retiredEpoch, token.Epoch) + + // Retirement is idempotent and does not advance the epoch again. + require.NoError(t, cache.RetireBucket(ctx, bucket)) + retiredEpochAgain, err := cache.rdb.Get(ctx, schedulerBucketKey(schedulerEpochPrefix, bucket)).Int64() + require.NoError(t, err) + require.Equal(t, retiredEpoch, retiredEpochAgain) + + // New readers miss because ready/active were removed atomically. A reader that + // captured activeVersion before retirement may still finish against that version. + _, hit, err := cache.GetSnapshot(ctx, bucket) + require.NoError(t, err) + require.False(t, hit) + ids, err := cache.rdb.ZRange(ctx, schedulerSnapshotKey(bucket, activeVersion), 0, -1).Result() + require.NoError(t, err) + require.Equal(t, []string{"4101"}, ids) + ttl, err := cache.rdb.TTL(ctx, schedulerSnapshotKey(bucket, activeVersion)).Result() + require.NoError(t, err) + require.Positive(t, ttl) + require.LessOrEqual(t, ttl, time.Duration(snapshotGraceTTLSeconds)*time.Second) + + buckets, err := cache.ListBuckets(ctx) + require.NoError(t, err) + require.NotContains(t, buckets, bucket) + require.Contains(t, buckets, otherBucket) + otherSnapshot, otherHit, err := cache.GetSnapshot(ctx, otherBucket) + require.NoError(t, err) + require.True(t, otherHit) + require.Len(t, otherSnapshot, 1) + require.Equal(t, otherAccount.ID, otherSnapshot[0].ID) + otherEpochAfter, err := cache.rdb.Get(ctx, schedulerBucketKey(schedulerEpochPrefix, otherBucket)).Int64() + require.NoError(t, err) + require.Equal(t, otherEpoch, otherEpochAfter) + + _, err = cache.CaptureBucketWriteToken(ctx, bucket) + require.ErrorIs(t, err, service.ErrSchedulerBucketRetired) + versionBeforeRejectedWrite, err := cache.rdb.Get(ctx, schedulerBucketKey(schedulerVersionPrefix, bucket)).Int64() + require.NoError(t, err) + err = cache.SetSnapshot(ctx, bucket, token, []service.Account{account}) + require.ErrorIs(t, err, service.ErrSchedulerBucketRetired) + versionAfterRejectedWrite, err := cache.rdb.Get(ctx, schedulerBucketKey(schedulerVersionPrefix, bucket)).Int64() + require.NoError(t, err) + require.Equal(t, versionBeforeRejectedWrite, versionAfterRejectedWrite, "fenced writers must not allocate a new version") + retired, err := cache.rdb.Exists(ctx, schedulerBucketKey(schedulerRetiredPrefix, bucket)).Result() + require.NoError(t, err) + require.EqualValues(t, 1, retired, "ordinary writers must never clear the tombstone") + mr.FastForward(time.Duration(snapshotGraceTTLSeconds+1) * time.Second) + exists, err := cache.rdb.Exists(ctx, schedulerSnapshotKey(bucket, activeVersion)).Result() + require.NoError(t, err) + require.Zero(t, exists, "retired active snapshot must expire after the in-flight grace period") + + newToken, err := cache.ReopenBucket(ctx, bucket) + require.NoError(t, err) + require.True(t, newToken.ValidFor(bucket)) + require.Equal(t, retiredEpoch, newToken.Epoch) + reopenedAgain, err := cache.ReopenBucket(ctx, bucket) + require.NoError(t, err) + require.Equal(t, newToken, reopenedAgain, "reopen must be idempotent within one retirement generation") + err = cache.SetSnapshot(ctx, bucket, token, []service.Account{account}) + require.ErrorIs(t, err, service.ErrSchedulerBucketWriteFenced) + require.NoError(t, cache.SetSnapshot(ctx, bucket, newToken, []service.Account{account})) + reopenedWhileOpen, err := cache.ReopenBucket(ctx, bucket) + require.NoError(t, err) + require.Equal(t, newToken, reopenedWhileOpen) + + snapshot, hit, err := cache.GetSnapshot(ctx, bucket) + require.NoError(t, err) + require.True(t, hit) + require.Len(t, snapshot, 1) + require.Equal(t, account.ID, snapshot[0].ID) +} + +func TestSchedulerCacheActivationIsFencedAfterRetire(t *testing.T) { + ctx := context.Background() + cache := newSchedulerCacheUnit(t) + bucket := service.SchedulerBucket{GroupID: 51, Platform: service.PlatformAnthropic, Mode: service.SchedulerModeMixed} + account := service.Account{ID: 5101, Platform: service.PlatformAnthropic, Type: service.AccountTypeAPIKey} + + token, err := cache.CaptureBucketWriteToken(ctx, bucket) + require.NoError(t, err) + version, err := cache.allocateSnapshotVersion(ctx, bucket, token) + require.NoError(t, err) + require.NoError(t, cache.writeSnapshotVersion(ctx, bucket, version, []service.Account{account})) + + // Deterministic race C: retirement and authoritative reopen both happen after + // INCR/write but before the old writer activates. + require.NoError(t, cache.RetireBucket(ctx, bucket)) + _, err = cache.ReopenBucket(ctx, bucket) + require.NoError(t, err) + err = cache.activateSnapshotVersion(ctx, bucket, token, version) + require.ErrorIs(t, err, service.ErrSchedulerBucketWriteFenced) + + exists, err := cache.rdb.Exists(ctx, schedulerSnapshotKey(bucket, version)).Result() + require.NoError(t, err) + require.Zero(t, exists, "fenced activation must delete its unpublished snapshot") + exists, err = cache.rdb.Exists( + ctx, + schedulerBucketKey(schedulerReadyPrefix, bucket), + schedulerBucketKey(schedulerActivePrefix, bucket), + ).Result() + require.NoError(t, err) + require.Zero(t, exists) + buckets, err := cache.ListBuckets(ctx) + require.NoError(t, err) + require.NotContains(t, buckets, bucket) +} + +func TestSchedulerCacheConcurrentReopenReturnsSameToken(t *testing.T) { + ctx := context.Background() + cache := newSchedulerCacheUnit(t) + bucket := service.SchedulerBucket{GroupID: 53, Platform: service.PlatformOpenAI, Mode: service.SchedulerModeForced} + account := service.Account{ID: 5301, Platform: service.PlatformOpenAI, Type: service.AccountTypeAPIKey} + + oldToken, err := cache.CaptureBucketWriteToken(ctx, bucket) + require.NoError(t, err) + require.NoError(t, cache.RetireBucket(ctx, bucket)) + + type reopenResult struct { + token service.SchedulerBucketWriteToken + err error + } + start := make(chan struct{}) + results := make(chan reopenResult, 2) + for range 2 { + go func() { + <-start + token, err := cache.ReopenBucket(ctx, bucket) + results <- reopenResult{token: token, err: err} + }() + } + close(start) + first := <-results + second := <-results + require.NoError(t, first.err) + require.NoError(t, second.err) + require.Equal(t, first.token, second.token) + require.Greater(t, first.token.Epoch, oldToken.Epoch) + + require.ErrorIs(t, cache.SetSnapshot(ctx, bucket, oldToken, []service.Account{account}), service.ErrSchedulerBucketWriteFenced) + require.NoError(t, cache.SetSnapshot(ctx, bucket, first.token, []service.Account{account})) +} + +func TestSchedulerCacheReopenExpiresPreviousActiveSnapshot(t *testing.T) { + ctx := context.Background() + cache, mr := newSchedulerCacheUnitWithRedis(t) + bucket := service.SchedulerBucket{GroupID: 52, Platform: service.PlatformGemini, Mode: service.SchedulerModeForced} + account := service.Account{ID: 5201, Platform: service.PlatformGemini, Type: service.AccountTypeAPIKey} + + oldToken, err := cache.CaptureBucketWriteToken(ctx, bucket) + require.NoError(t, err) + require.NoError(t, cache.SetSnapshot(ctx, bucket, oldToken, []service.Account{account})) + oldVersion, err := cache.rdb.Get(ctx, schedulerBucketKey(schedulerActivePrefix, bucket)).Result() + require.NoError(t, err) + retiredEpoch := oldToken.Epoch + 1 + require.NoError(t, cache.rdb.Set(ctx, schedulerBucketKey(schedulerEpochPrefix, bucket), retiredEpoch, 0).Err()) + require.NoError(t, cache.rdb.Set(ctx, schedulerBucketKey(schedulerRetiredPrefix, bucket), retiredEpoch, 0).Err()) + + newToken, err := cache.ReopenBucket(ctx, bucket) + require.NoError(t, err) + require.Equal(t, retiredEpoch, newToken.Epoch) + _, hit, err := cache.GetSnapshot(ctx, bucket) + require.NoError(t, err) + require.False(t, hit) + ttl, err := cache.rdb.TTL(ctx, schedulerSnapshotKey(bucket, oldVersion)).Result() + require.NoError(t, err) + require.Positive(t, ttl) + require.LessOrEqual(t, ttl, time.Duration(snapshotGraceTTLSeconds)*time.Second) + + require.ErrorIs(t, cache.SetSnapshot(ctx, bucket, oldToken, []service.Account{account}), service.ErrSchedulerBucketWriteFenced) + mr.FastForward(time.Duration(snapshotGraceTTLSeconds+1) * time.Second) + exists, err := cache.rdb.Exists(ctx, schedulerSnapshotKey(bucket, oldVersion)).Result() + require.NoError(t, err) + require.Zero(t, exists) + require.NoError(t, cache.SetSnapshot(ctx, bucket, newToken, []service.Account{account})) +} diff --git a/backend/internal/service/scheduler_cache.go b/backend/internal/service/scheduler_cache.go index f9794c8214..b117c6194f 100644 --- a/backend/internal/service/scheduler_cache.go +++ b/backend/internal/service/scheduler_cache.go @@ -2,6 +2,7 @@ package service import ( "context" + "errors" "fmt" "strconv" "strings" @@ -14,6 +15,22 @@ const ( SchedulerModeForced = "forced" ) +var ( + ErrSchedulerBucketRetired = errors.New("scheduler bucket retired") + ErrSchedulerBucketWriteFenced = errors.New("scheduler bucket write fenced") +) + +// SchedulerBucketWriteToken fences a snapshot writer to one bucket epoch. +// Tokens must be captured before any database load or queued rebuild work. +type SchedulerBucketWriteToken struct { + Bucket SchedulerBucket + Epoch int64 +} + +func (t SchedulerBucketWriteToken) ValidFor(bucket SchedulerBucket) bool { + return t.Epoch > 0 && t.Bucket == bucket +} + type SchedulerBucket struct { GroupID int64 Platform string @@ -47,8 +64,21 @@ func ParseSchedulerBucket(raw string) (SchedulerBucket, bool) { type SchedulerCache interface { // GetSnapshot 读取快照并返回命中与否(ready + active + 数据完整)。 GetSnapshot(ctx context.Context, bucket SchedulerBucket) ([]*Account, bool, error) - // SetSnapshot 写入快照并切换激活版本。 - SetSnapshot(ctx context.Context, bucket SchedulerBucket, accounts []Account) error + // CaptureBucketWriteToken captures the current open epoch without changing + // retirement state. A tombstoned bucket returns ErrSchedulerBucketRetired. + CaptureBucketWriteToken(ctx context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) + // SetSnapshot 写入快照并切换激活版本。token 必须在 DB load/任务排队前取得。 + SetSnapshot(ctx context.Context, bucket SchedulerBucket, token SchedulerBucketWriteToken, accounts []Account) error + // RetireBucket persistently tombstones a bucket and fences every older writer. + // Readers that captured the active version before retirement may finish; new + // readers observe ready/active as absent. + RetireBucket(ctx context.Context, bucket SchedulerBucket) error + // ReopenBucket is the only operation allowed to clear a tombstone. It returns + // the retirement generation established by RetireBucket; repeated calls for + // the same generation are idempotent. Callers must serialize a fresh authority + // check through ReopenBucket with RetireBucket under the same bucket lifecycle + // lock; ordinary rebuild paths never call ReopenBucket. + ReopenBucket(ctx context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) // GetAccount 获取单账号快照。 GetAccount(ctx context.Context, accountID int64) (*Account, error) // SetAccount 写入单账号快照(包含不可调度状态)。 diff --git a/backend/internal/service/scheduler_snapshot_full_rebuild_test.go b/backend/internal/service/scheduler_snapshot_full_rebuild_test.go index 09ec2790ce..7ef4d9ed9f 100644 --- a/backend/internal/service/scheduler_snapshot_full_rebuild_test.go +++ b/backend/internal/service/scheduler_snapshot_full_rebuild_test.go @@ -34,6 +34,14 @@ func (c *schedulerFullRebuildTestCache) TryLockBucket(context.Context, Scheduler return false, nil } +func (c *schedulerFullRebuildTestCache) CaptureBucketWriteToken(_ context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) { + return SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} + +func (c *schedulerFullRebuildTestCache) ReopenBucket(_ context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) { + return SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} + func TestSchedulerSnapshotServiceFullRebuildCoalescesConcurrentRequestsIntoTrailingRun(t *testing.T) { svc := &SchedulerSnapshotService{} wantTrailingErr := errors.New("trailing rebuild failed") diff --git a/backend/internal/service/scheduler_snapshot_hydration_test.go b/backend/internal/service/scheduler_snapshot_hydration_test.go index 0a1d0a0a52..5238529c63 100644 --- a/backend/internal/service/scheduler_snapshot_hydration_test.go +++ b/backend/internal/service/scheduler_snapshot_hydration_test.go @@ -19,10 +19,22 @@ func (c *snapshotHydrationCache) GetSnapshot(ctx context.Context, bucket Schedul return c.snapshot, true, nil } -func (c *snapshotHydrationCache) SetSnapshot(ctx context.Context, bucket SchedulerBucket, accounts []Account) error { +func (c *snapshotHydrationCache) CaptureBucketWriteToken(ctx context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) { + return SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} + +func (c *snapshotHydrationCache) SetSnapshot(ctx context.Context, bucket SchedulerBucket, token SchedulerBucketWriteToken, accounts []Account) error { return nil } +func (c *snapshotHydrationCache) RetireBucket(ctx context.Context, bucket SchedulerBucket) error { + return nil +} + +func (c *snapshotHydrationCache) ReopenBucket(ctx context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) { + return SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} + func (c *snapshotHydrationCache) GetAccount(ctx context.Context, accountID int64) (*Account, error) { if c.accounts == nil { return nil, nil diff --git a/backend/internal/service/scheduler_snapshot_outbox_cleanup_test.go b/backend/internal/service/scheduler_snapshot_outbox_cleanup_test.go index 91e2d36f76..acbdcdb259 100644 --- a/backend/internal/service/scheduler_snapshot_outbox_cleanup_test.go +++ b/backend/internal/service/scheduler_snapshot_outbox_cleanup_test.go @@ -21,10 +21,22 @@ func (c *outboxCleanupCache) GetSnapshot(ctx context.Context, bucket SchedulerBu return nil, false, nil } -func (c *outboxCleanupCache) SetSnapshot(ctx context.Context, bucket SchedulerBucket, accounts []Account) error { +func (c *outboxCleanupCache) CaptureBucketWriteToken(ctx context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) { + return SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} + +func (c *outboxCleanupCache) SetSnapshot(ctx context.Context, bucket SchedulerBucket, token SchedulerBucketWriteToken, accounts []Account) error { return nil } +func (c *outboxCleanupCache) RetireBucket(ctx context.Context, bucket SchedulerBucket) error { + return nil +} + +func (c *outboxCleanupCache) ReopenBucket(ctx context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) { + return SchedulerBucketWriteToken{Bucket: bucket, Epoch: 1}, nil +} + func (c *outboxCleanupCache) GetAccount(ctx context.Context, accountID int64) (*Account, error) { return nil, nil } diff --git a/backend/internal/service/scheduler_snapshot_retirement_test.go b/backend/internal/service/scheduler_snapshot_retirement_test.go new file mode 100644 index 0000000000..e1ac974f81 --- /dev/null +++ b/backend/internal/service/scheduler_snapshot_retirement_test.go @@ -0,0 +1,289 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +type retirementRaceCache struct { + SchedulerCache + + mu sync.Mutex + epochs map[string]int64 + retired map[string]bool + listBuckets []SchedulerBucket + captures []SchedulerBucket + reopens []SchedulerBucket + setAttempts map[string]int + published map[string]int + versions map[string]int + beforeSet func() +} + +func newRetirementRaceCache(buckets ...SchedulerBucket) *retirementRaceCache { + return &retirementRaceCache{ + epochs: make(map[string]int64), + retired: make(map[string]bool), + listBuckets: buckets, + setAttempts: make(map[string]int), + published: make(map[string]int), + versions: make(map[string]int), + } +} + +func (c *retirementRaceCache) GetSnapshot(context.Context, SchedulerBucket) ([]*Account, bool, error) { + return nil, false, nil +} + +func (c *retirementRaceCache) CaptureBucketWriteToken(_ context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) { + c.mu.Lock() + defer c.mu.Unlock() + key := bucket.String() + c.captures = append(c.captures, bucket) + if c.retired[key] { + return SchedulerBucketWriteToken{}, ErrSchedulerBucketRetired + } + if c.epochs[key] == 0 { + c.epochs[key] = 1 + } + return SchedulerBucketWriteToken{Bucket: bucket, Epoch: c.epochs[key]}, nil +} + +func (c *retirementRaceCache) SetSnapshot(_ context.Context, bucket SchedulerBucket, token SchedulerBucketWriteToken, _ []Account) error { + if c.beforeSet != nil { + c.beforeSet() + } + c.mu.Lock() + defer c.mu.Unlock() + key := bucket.String() + c.setAttempts[key]++ + if !token.ValidFor(bucket) { + return ErrSchedulerBucketWriteFenced + } + if c.retired[key] { + return ErrSchedulerBucketRetired + } + if c.epochs[key] != token.Epoch { + return ErrSchedulerBucketWriteFenced + } + c.versions[key]++ + c.published[key]++ + return nil +} + +func (c *retirementRaceCache) RetireBucket(_ context.Context, bucket SchedulerBucket) error { + c.mu.Lock() + defer c.mu.Unlock() + key := bucket.String() + if !c.retired[key] { + c.epochs[key]++ + if c.epochs[key] < 1 { + c.epochs[key] = 1 + } + c.retired[key] = true + } + return nil +} + +func (c *retirementRaceCache) ReopenBucket(_ context.Context, bucket SchedulerBucket) (SchedulerBucketWriteToken, error) { + c.mu.Lock() + defer c.mu.Unlock() + key := bucket.String() + if c.epochs[key] == 0 { + c.epochs[key] = 1 + } + delete(c.retired, key) + c.reopens = append(c.reopens, bucket) + return SchedulerBucketWriteToken{Bucket: bucket, Epoch: c.epochs[key]}, nil +} + +func (c *retirementRaceCache) TryLockBucket(context.Context, SchedulerBucket, time.Duration) (bool, error) { + return true, nil +} + +func (c *retirementRaceCache) UnlockBucket(context.Context, SchedulerBucket) error { + return nil +} + +func (c *retirementRaceCache) ListBuckets(context.Context) ([]SchedulerBucket, error) { + return append([]SchedulerBucket(nil), c.listBuckets...), nil +} + +func (c *retirementRaceCache) counts(bucket SchedulerBucket) (setAttempts, published int) { + c.mu.Lock() + defer c.mu.Unlock() + return c.setAttempts[bucket.String()], c.published[bucket.String()] +} + +func (c *retirementRaceCache) captureAndReopenCounts() (captures, reopens int) { + c.mu.Lock() + defer c.mu.Unlock() + return len(c.captures), len(c.reopens) +} + +func (c *retirementRaceCache) version(bucket SchedulerBucket) int { + c.mu.Lock() + defer c.mu.Unlock() + return c.versions[bucket.String()] +} + +type retirementGroupRepo struct { + GroupRepository + groups []Group + err error +} + +func (r *retirementGroupRepo) ListActive(context.Context) ([]Group, error) { + return r.groups, r.err +} + +func TestSchedulerFullRebuildCapturesAllRegistryTokensBeforeDBLoad(t *testing.T) { + first := SchedulerBucket{GroupID: 61, Platform: PlatformOpenAI, Mode: SchedulerModeSingle} + queued := SchedulerBucket{GroupID: 61, Platform: PlatformOpenAI, Mode: SchedulerModeForced} + cache := newRetirementRaceCache(first, queued) + dbStarted := make(chan struct{}) + releaseDB := make(chan struct{}) + var firstDB sync.Once + repo := &mockAccountRepoForPlatform{ + listPlatformFunc: func(context.Context, string) ([]Account, error) { + firstDB.Do(func() { + close(dbStarted) + <-releaseDB + }) + return []Account{{ID: 6101, Platform: PlatformOpenAI, Status: StatusActive, Schedulable: true}}, nil + }, + } + svc := NewSchedulerSnapshotService(cache, nil, repo, nil, &config.Config{ + RunMode: config.RunModeStandard, + Gateway: config.GatewayConfig{Scheduling: config.GatewaySchedulingConfig{ + DbFallbackEnabled: true, + }}, + }) + + result := make(chan error, 1) + go func() { result <- svc.triggerFullRebuild("retirement_race_a") }() + select { + case <-dbStarted: + case <-time.After(time.Second): + t.Fatal("first DB load did not start") + } + + captures, reopens := cache.captureAndReopenCounts() + require.Equal(t, 2, captures, "all registry tokens must be captured before the first DB load") + require.Zero(t, reopens) + require.NoError(t, cache.RetireBucket(context.Background(), queued)) + _, err := cache.ReopenBucket(context.Background(), queued) + require.NoError(t, err) + close(releaseDB) + require.NoError(t, <-result) + + _, firstPublished := cache.counts(first) + queuedAttempts, queuedPublished := cache.counts(queued) + require.Equal(t, 1, firstPublished) + require.Equal(t, 1, queuedAttempts) + require.Zero(t, queuedPublished, "queued registry task must not adopt the reopened epoch") +} + +func TestSchedulerRebuildRetireAfterDBLoadFencesPublish(t *testing.T) { + bucket := SchedulerBucket{GroupID: 62, Platform: PlatformOpenAI, Mode: SchedulerModeSingle} + cache := newRetirementRaceCache() + dbReturned := make(chan struct{}) + setEntered := make(chan struct{}) + releaseSet := make(chan struct{}) + cache.beforeSet = func() { + close(setEntered) + <-releaseSet + } + repo := &mockAccountRepoForPlatform{ + listPlatformFunc: func(context.Context, string) ([]Account, error) { + close(dbReturned) + return []Account{{ID: 6201, Platform: PlatformOpenAI, Status: StatusActive, Schedulable: true}}, nil + }, + } + svc := NewSchedulerSnapshotService(cache, nil, repo, nil, &config.Config{ + RunMode: config.RunModeStandard, + Gateway: config.GatewayConfig{Scheduling: config.GatewaySchedulingConfig{ + DbFallbackEnabled: true, + }}, + }) + + result := make(chan error, 1) + go func() { + result <- svc.rebuildBuckets(context.Background(), []SchedulerBucket{bucket}, "retirement_race_b") + }() + select { + case <-dbReturned: + case <-time.After(time.Second): + t.Fatal("DB load did not return") + } + select { + case <-setEntered: + case <-time.After(time.Second): + t.Fatal("snapshot writer did not reach allocation boundary") + } + require.NoError(t, cache.RetireBucket(context.Background(), bucket)) + close(releaseSet) + require.NoError(t, <-result) + + setAttempts, published := cache.counts(bucket) + require.Equal(t, 1, setAttempts) + require.Zero(t, published) + require.Zero(t, cache.version(bucket), "retirement before allocation must not advance the snapshot version") +} + +func TestSchedulerFallbackReturnsDBAccountsWhenBucketRetired(t *testing.T) { + bucket := SchedulerBucket{GroupID: 63, Platform: PlatformOpenAI, Mode: SchedulerModeSingle} + cache := newRetirementRaceCache() + require.NoError(t, cache.RetireBucket(context.Background(), bucket)) + repo := &mockAccountRepoForPlatform{ + accounts: []Account{{ID: 6301, Platform: PlatformOpenAI, Status: StatusActive, Schedulable: true}}, + } + svc := NewSchedulerSnapshotService(cache, nil, repo, nil, &config.Config{ + RunMode: config.RunModeStandard, + Gateway: config.GatewayConfig{Scheduling: config.GatewaySchedulingConfig{ + DbFallbackEnabled: true, + }}, + }) + groupID := bucket.GroupID + + accounts, useMixed, err := svc.ListSchedulableAccounts(context.Background(), &groupID, bucket.Platform, false) + require.NoError(t, err) + require.False(t, useMixed) + require.Len(t, accounts, 1) + setAttempts, published := cache.counts(bucket) + require.Zero(t, setAttempts) + require.Zero(t, published) +} + +func TestSchedulerDefaultBucketsUseCaptureAndListActiveFailureKeepsGroupZero(t *testing.T) { + cache := newRetirementRaceCache() + svc := NewSchedulerSnapshotService( + cache, + nil, + nil, + &retirementGroupRepo{err: errors.New("list active failed")}, + testConfig(), + ) + + buckets, err := svc.defaultBuckets(context.Background()) + require.NoError(t, err) + require.NotEmpty(t, buckets) + for _, bucket := range buckets { + require.Zero(t, bucket.GroupID) + } + + tasks, err := svc.prepareBucketWriteTasks(context.Background(), buckets) + require.NoError(t, err) + require.Len(t, tasks, len(buckets)) + captures, reopens := cache.captureAndReopenCounts() + require.Equal(t, len(buckets), captures) + require.Zero(t, reopens) +} diff --git a/backend/internal/service/scheduler_snapshot_service.go b/backend/internal/service/scheduler_snapshot_service.go index 77bef26edf..a40e42fc19 100644 --- a/backend/internal/service/scheduler_snapshot_service.go +++ b/backend/internal/service/scheduler_snapshot_service.go @@ -31,6 +31,11 @@ type batchSeenKey struct { platform string } +type schedulerBucketWriteTask struct { + bucket SchedulerBucket + token SchedulerBucketWriteToken +} + type SchedulerSnapshotService struct { cache SchedulerCache outboxRepo SchedulerOutboxRepository @@ -117,6 +122,8 @@ func (s *SchedulerSnapshotService) ListSchedulableAccounts(ctx context.Context, useMixed := (platform == PlatformAnthropic || platform == PlatformGemini) && !hasForcePlatform mode := s.resolveMode(platform, hasForcePlatform) bucket := s.bucketFor(groupID, platform, mode) + var writeToken SchedulerBucketWriteToken + canPublish := false if s.cache != nil { cached, hit, err := s.cache.GetSnapshot(ctx, bucket) @@ -125,6 +132,17 @@ func (s *SchedulerSnapshotService) ListSchedulableAccounts(ctx context.Context, } else if hit { return derefAccounts(cached), useMixed, nil } + token, err := s.cache.CaptureBucketWriteToken(ctx, bucket) + if err != nil { + if errors.Is(err, ErrSchedulerBucketRetired) || errors.Is(err, ErrSchedulerBucketWriteFenced) { + slog.Debug("[Scheduler] cache publish fenced", "bucket", bucket.String()) + } else { + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] cache publish token failed: bucket=%s err=%v", bucket.String(), err) + } + } else { + writeToken = token + canPublish = true + } } if err := s.guardFallback(ctx); err != nil { @@ -139,9 +157,13 @@ func (s *SchedulerSnapshotService) ListSchedulableAccounts(ctx context.Context, return nil, useMixed, err } - if s.cache != nil { - if err := s.cache.SetSnapshot(fallbackCtx, bucket, accounts); err != nil { - logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] cache write failed: bucket=%s err=%v", bucket.String(), err) + if s.cache != nil && canPublish { + if err := s.cache.SetSnapshot(fallbackCtx, bucket, writeToken, accounts); err != nil { + if errors.Is(err, ErrSchedulerBucketRetired) || errors.Is(err, ErrSchedulerBucketWriteFenced) { + slog.Debug("[Scheduler] cache publish fenced", "bucket", bucket.String()) + } else { + logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] cache write failed: bucket=%s err=%v", bucket.String(), err) + } } } @@ -507,19 +529,12 @@ func (s *SchedulerSnapshotService) rebuildByAccount(ctx context.Context, account return nil } - var firstErr error - if err := s.rebuildBucketsForPlatform(ctx, account.Platform, groupIDs, reason, seen); err != nil && firstErr == nil { - firstErr = err - } + buckets := s.bucketsForPlatform(account.Platform, groupIDs, seen) if account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled() { - if err := s.rebuildBucketsForPlatform(ctx, PlatformAnthropic, groupIDs, reason, seen); err != nil && firstErr == nil { - firstErr = err - } - if err := s.rebuildBucketsForPlatform(ctx, PlatformGemini, groupIDs, reason, seen); err != nil && firstErr == nil { - firstErr = err - } + buckets = append(buckets, s.bucketsForPlatform(PlatformAnthropic, groupIDs, seen)...) + buckets = append(buckets, s.bucketsForPlatform(PlatformGemini, groupIDs, seen)...) } - return firstErr + return s.rebuildBuckets(ctx, buckets, reason) } func (s *SchedulerSnapshotService) rebuildByGroupIDs(ctx context.Context, groupIDs []int64, reason string, seen map[batchSeenKey]struct{}) error { @@ -528,20 +543,18 @@ func (s *SchedulerSnapshotService) rebuildByGroupIDs(ctx context.Context, groupI return nil } platforms := []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok} - var firstErr error + buckets := make([]SchedulerBucket, 0, len(groupIDs)*12) for _, platform := range platforms { - if err := s.rebuildBucketsForPlatform(ctx, platform, groupIDs, reason, seen); err != nil && firstErr == nil { - firstErr = err - } + buckets = append(buckets, s.bucketsForPlatform(platform, groupIDs, seen)...) } - return firstErr + return s.rebuildBuckets(ctx, buckets, reason) } -func (s *SchedulerSnapshotService) rebuildBucketsForPlatform(ctx context.Context, platform string, groupIDs []int64, reason string, seen map[batchSeenKey]struct{}) error { +func (s *SchedulerSnapshotService) bucketsForPlatform(platform string, groupIDs []int64, seen map[batchSeenKey]struct{}) []SchedulerBucket { if platform == "" { return nil } - var firstErr error + buckets := make([]SchedulerBucket, 0, len(groupIDs)*3) for _, gid := range groupIDs { // Within a single poll batch, skip (groupID, platform) pairs that were // already rebuilt. The first rebuild loads fresh DB data for all accounts @@ -554,35 +567,52 @@ func (s *SchedulerSnapshotService) rebuildBucketsForPlatform(ctx context.Context } seen[key] = struct{}{} } - if err := s.rebuildBucket(ctx, SchedulerBucket{GroupID: gid, Platform: platform, Mode: SchedulerModeSingle}, reason); err != nil && firstErr == nil { - firstErr = err - } - if err := s.rebuildBucket(ctx, SchedulerBucket{GroupID: gid, Platform: platform, Mode: SchedulerModeForced}, reason); err != nil && firstErr == nil { - firstErr = err - } + buckets = append(buckets, SchedulerBucket{GroupID: gid, Platform: platform, Mode: SchedulerModeSingle}) + buckets = append(buckets, SchedulerBucket{GroupID: gid, Platform: platform, Mode: SchedulerModeForced}) if platform == PlatformAnthropic || platform == PlatformGemini { - if err := s.rebuildBucket(ctx, SchedulerBucket{GroupID: gid, Platform: platform, Mode: SchedulerModeMixed}, reason); err != nil && firstErr == nil { - firstErr = err - } + buckets = append(buckets, SchedulerBucket{GroupID: gid, Platform: platform, Mode: SchedulerModeMixed}) } } - return firstErr + return buckets } func (s *SchedulerSnapshotService) rebuildBuckets(ctx context.Context, buckets []SchedulerBucket, reason string) error { - var firstErr error - for _, bucket := range buckets { - if err := s.rebuildBucket(ctx, bucket, reason); err != nil && firstErr == nil { + tasks, firstErr := s.prepareBucketWriteTasks(ctx, buckets) + for _, task := range tasks { + if err := s.rebuildBucketWithToken(ctx, task, reason); err != nil && firstErr == nil { firstErr = err } } return firstErr } -func (s *SchedulerSnapshotService) rebuildBucket(ctx context.Context, bucket SchedulerBucket, reason string) error { +func (s *SchedulerSnapshotService) prepareBucketWriteTasks(ctx context.Context, buckets []SchedulerBucket) ([]schedulerBucketWriteTask, error) { + if s.cache == nil { + return nil, ErrSchedulerCacheNotReady + } + tasks := make([]schedulerBucketWriteTask, 0, len(buckets)) + var firstErr error + for _, bucket := range buckets { + token, err := s.cache.CaptureBucketWriteToken(ctx, bucket) + if err != nil { + if errors.Is(err, ErrSchedulerBucketRetired) || errors.Is(err, ErrSchedulerBucketWriteFenced) { + continue + } + if firstErr == nil { + firstErr = err + } + continue + } + tasks = append(tasks, schedulerBucketWriteTask{bucket: bucket, token: token}) + } + return tasks, firstErr +} + +func (s *SchedulerSnapshotService) rebuildBucketWithToken(ctx context.Context, task schedulerBucketWriteTask, reason string) error { if s.cache == nil { return ErrSchedulerCacheNotReady } + bucket := task.bucket ok, err := s.cache.TryLockBucket(ctx, bucket, 30*time.Second) if err != nil { return err @@ -602,7 +632,11 @@ func (s *SchedulerSnapshotService) rebuildBucket(ctx context.Context, bucket Sch logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] rebuild failed: bucket=%s reason=%s err=%v", bucket.String(), reason, err) return err } - if err := s.cache.SetSnapshot(rebuildCtx, bucket, accounts); err != nil { + if err := s.cache.SetSnapshot(rebuildCtx, bucket, task.token, accounts); err != nil { + if errors.Is(err, ErrSchedulerBucketRetired) || errors.Is(err, ErrSchedulerBucketWriteFenced) { + slog.Debug("[Scheduler] rebuild fenced", "bucket", bucket.String(), "reason", reason) + return nil + } logger.LegacyPrintf("service.scheduler_snapshot", "[Scheduler] rebuild cache failed: bucket=%s reason=%s err=%v", bucket.String(), reason, err) return err }