diff --git a/backend/cmd/server/wire.go b/backend/cmd/server/wire.go index 9bfa2717..b474cfa1 100644 --- a/backend/cmd/server/wire.go +++ b/backend/cmd/server/wire.go @@ -98,6 +98,7 @@ func provideCleanup( backupSvc *service.BackupService, paymentOrderExpiry *service.PaymentOrderExpiryService, channelMonitorRunner *service.ChannelMonitorRunner, + quotaFlusher *service.UserPlatformQuotaUsageFlusher, ) func() { return func() { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -246,6 +247,12 @@ func provideCleanup( } return nil }}, + {"UserPlatformQuotaUsageFlusher", func() error { + if quotaFlusher != nil { + quotaFlusher.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 6e8be8fc..9b059c5d 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -269,7 +269,8 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { scheduledTestRunnerService := service.ProvideScheduledTestRunnerService(scheduledTestPlanRepository, scheduledTestService, accountTestService, rateLimitService, configConfig) paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService) channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService) - v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner) + userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService) + v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher) application := &Application{ Server: httpServer, Cleanup: v, @@ -324,6 +325,7 @@ func provideCleanup( backupSvc *service.BackupService, paymentOrderExpiry *service.PaymentOrderExpiryService, channelMonitorRunner *service.ChannelMonitorRunner, + quotaFlusher *service.UserPlatformQuotaUsageFlusher, ) func() { return func() { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -471,6 +473,12 @@ func provideCleanup( } return nil }}, + {"UserPlatformQuotaUsageFlusher", func() error { + if quotaFlusher != nil { + quotaFlusher.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index a44b2d5c..7f4e4773 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -77,6 +77,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { nil, // backupSvc nil, // paymentOrderExpiry nil, // channelMonitorRunner + nil, // quotaFlusher ) require.NotPanics(t, func() { diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index dcbf30b4..df9dcefc 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -1094,6 +1094,13 @@ type DatabaseConfig struct { ConnMaxLifetimeMinutes int `mapstructure:"conn_max_lifetime_minutes"` // ConnMaxIdleTimeMinutes: 空闲连接最大存活时间,及时释放不活跃连接 ConnMaxIdleTimeMinutes int `mapstructure:"conn_max_idle_time_minutes"` + // UserPlatformQuotaFlusherEnabled: 是否启用 user×platform 配额写聚合 flusher + UserPlatformQuotaFlusherEnabled bool `mapstructure:"user_platform_quota_flusher_enabled"` + // UserPlatformQuotaFlushIntervalMs: flusher 刷写间隔(毫秒) + UserPlatformQuotaFlushIntervalMs int `mapstructure:"user_platform_quota_flush_interval_ms"` + // UserPlatformQuotaFlushBatchSize: flusher 单批最大条数 + // 建议 ≤ 6000(单条 UPSERT 原子上限) + UserPlatformQuotaFlushBatchSize int `mapstructure:"user_platform_quota_flush_batch_size"` } func (d *DatabaseConfig) DSN() string { @@ -1661,6 +1668,9 @@ func setDefaults() { viper.SetDefault("database.max_idle_conns", 128) viper.SetDefault("database.conn_max_lifetime_minutes", 30) viper.SetDefault("database.conn_max_idle_time_minutes", 5) + viper.SetDefault("database.user_platform_quota_flusher_enabled", false) + viper.SetDefault("database.user_platform_quota_flush_interval_ms", 2000) + viper.SetDefault("database.user_platform_quota_flush_batch_size", 1000) // Redis viper.SetDefault("redis.host", "localhost") diff --git a/backend/internal/handler/admin/user_handler.go b/backend/internal/handler/admin/user_handler.go index 32a21692..6c0a02ff 100644 --- a/backend/internal/handler/admin/user_handler.go +++ b/backend/internal/handler/admin/user_handler.go @@ -743,7 +743,7 @@ func (h *UserHandler) UpdateUserPlatformQuotas(c *gin.Context) { if h.billingCache != nil { for _, p := range service.AllowedQuotaPlatforms { if err := h.billingCache.DeleteUserPlatformQuotaCache(ctx, userID, p); err != nil { - slog.Warn("quota cache invalidation failed", "user_id", userID, "platform", p, "err", err) + slog.Error("ALERT: quota cache invalidation failed after UpsertForUser; limit 生效可能延迟至 sentinel TTL(最长 1h),需人工确认或重试失效", "user_id", userID, "platform", p, "err", err) } } } @@ -827,7 +827,7 @@ func (h *UserHandler) ResetUserPlatformQuotaWindow(c *gin.Context) { if h.billingCache != nil { if err := h.billingCache.DeleteUserPlatformQuotaCache(ctx, userID, req.Platform); err != nil { - slog.Warn("quota cache invalidation failed", "user_id", userID, "platform", req.Platform, "err", err) + slog.Error("ALERT: quota cache invalidation failed after ResetExpiredWindow; 窗口重置可能延迟至 sentinel TTL(最长 1h)", "user_id", userID, "platform", req.Platform, "err", err) } } diff --git a/backend/internal/repository/billing_cache.go b/backend/internal/repository/billing_cache.go index 60dae954..de229da9 100644 --- a/backend/internal/repository/billing_cache.go +++ b/backend/internal/repository/billing_cache.go @@ -7,6 +7,7 @@ import ( "log" "math/rand/v2" "strconv" + "strings" "time" "github.com/Wei-Shaw/sub2api/internal/service" @@ -338,38 +339,26 @@ func userPlatformQuotaCacheKey(userID int64, platform string) string { return fmt.Sprintf("billing:user_platform_quota:%d:%s", userID, platform) } -func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) (*service.UserPlatformQuotaCacheEntry, bool, error) { - key := userPlatformQuotaCacheKey(userID, platform) - fields := []string{ - "daily_usage", "weekly_usage", "monthly_usage", "version", "schema_version", - "daily_limit", "weekly_limit", "monthly_limit", - "daily_window_start", "weekly_window_start", "monthly_window_start", +// parseUserPlatformQuotaHash 将 Redis HGETALL 返回的 map[string]string 反序列化为 +// *service.UserPlatformQuotaCacheEntry。空 map(key 不存在)返回 nil。 +// GetUserPlatformQuotaCache 和 BatchGetUserPlatformQuotaCache 共用此函数,确保解析逻辑一致。 +func parseUserPlatformQuotaHash(m map[string]string) *service.UserPlatformQuotaCacheEntry { + if len(m) == 0 { + return nil } - vals, err := c.rdb.HMGet(ctx, key, fields...).Result() - if err != nil { - return nil, false, err - } - // 前4个全为nil → key 不存在 - if vals[0] == nil && vals[1] == nil && vals[2] == nil && vals[3] == nil { - return nil, false, nil - } - parseFloat := func(v any) float64 { - if v == nil { + parseFloat := func(s string) float64 { + if s == "" { return 0 } - s, ok := v.(string) - if !ok { + f, err := strconv.ParseFloat(s, 64) + if err != nil { + log.Printf("billing_cache: corrupt quota usage field %q (using 0): %v", s, err) return 0 } - f, _ := strconv.ParseFloat(s, 64) return f } - parseFloatPtr := func(v any) *float64 { - if v == nil { - return nil - } - s, ok := v.(string) - if !ok || s == "" { + parseFloatPtr := func(s string) *float64 { + if s == "" { return nil } f, err := strconv.ParseFloat(s, 64) @@ -378,12 +367,8 @@ func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int } return &f } - parseTimePtr := func(v any) *time.Time { - if v == nil { - return nil - } - s, ok := v.(string) - if !ok || s == "" { + parseTimePtr := func(s string) *time.Time { + if s == "" { return nil } n, err := strconv.ParseInt(s, 10, 64) @@ -393,30 +378,37 @@ func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int t := time.Unix(n, 0).UTC() return &t } - parseInt64 := func(v any) int64 { - if v == nil { - return 0 - } - s, ok := v.(string) - if !ok { - return 0 - } + parseInt64 := func(s string) int64 { n, _ := strconv.ParseInt(s, 10, 64) return n } return &service.UserPlatformQuotaCacheEntry{ - DailyUsageUSD: parseFloat(vals[0]), - WeeklyUsageUSD: parseFloat(vals[1]), - MonthlyUsageUSD: parseFloat(vals[2]), - Version: parseInt64(vals[3]), - SchemaVersion: parseInt64(vals[4]), - DailyLimitUSD: parseFloatPtr(vals[5]), - WeeklyLimitUSD: parseFloatPtr(vals[6]), - MonthlyLimitUSD: parseFloatPtr(vals[7]), - DailyWindowStart: parseTimePtr(vals[8]), - WeeklyWindowStart: parseTimePtr(vals[9]), - MonthlyWindowStart: parseTimePtr(vals[10]), - }, true, nil + DailyUsageUSD: parseFloat(m["daily_usage"]), + WeeklyUsageUSD: parseFloat(m["weekly_usage"]), + MonthlyUsageUSD: parseFloat(m["monthly_usage"]), + Version: parseInt64(m["version"]), + SchemaVersion: parseInt64(m["schema_version"]), + DailyLimitUSD: parseFloatPtr(m["daily_limit"]), + WeeklyLimitUSD: parseFloatPtr(m["weekly_limit"]), + MonthlyLimitUSD: parseFloatPtr(m["monthly_limit"]), + DailyWindowStart: parseTimePtr(m["daily_window_start"]), + WeeklyWindowStart: parseTimePtr(m["weekly_window_start"]), + MonthlyWindowStart: parseTimePtr(m["monthly_window_start"]), + } +} + +func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) (*service.UserPlatformQuotaCacheEntry, bool, error) { + key := userPlatformQuotaCacheKey(userID, platform) + m, err := c.rdb.HGetAll(ctx, key).Result() + if err != nil { + return nil, false, err + } + entry := parseUserPlatformQuotaHash(m) + if entry == nil { + // 空 map → key 不存在 → MISS + return nil, false, nil + } + return entry, true, nil } func (c *billingCache) SetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string, entry *service.UserPlatformQuotaCacheEntry, ttl time.Duration) error { @@ -468,9 +460,12 @@ func (c *billingCache) DeleteUserPlatformQuotaCache(ctx context.Context, userID // SetCache 重建为新版 entry —— 若此处仍累加,上层覆盖时会丢失这部分增量,导致 Redis usage 比真实偏小。 // key 不存在同样跳过(由下次 SetCache 重建)。 // KEYS[1] = hash key +// KEYS[2] = 脏集 key(dirty set) // ARGV[1] = cost (string float) // ARGV[2] = ttl seconds // ARGV[3] = expected schema_version (Go 侧 UserPlatformQuotaCacheSchemaV1) +// ARGV[4] = dirty set member(空串则不 SADD) +// ARGV[5] = 脏集兜底 TTL 秒 const updateUserPlatformQuotaUsageScript = ` if redis.call("EXISTS", KEYS[1]) == 0 then return 0 @@ -484,18 +479,125 @@ redis.call("HINCRBYFLOAT", KEYS[1], "weekly_usage", ARGV[1]) redis.call("HINCRBYFLOAT", KEYS[1], "monthly_usage", ARGV[1]) redis.call("HINCRBY", KEYS[1], "version", 1) redis.call("EXPIRE", KEYS[1], ARGV[2]) +if ARGV[4] ~= "" then + redis.call("SADD", KEYS[2], ARGV[4]) + redis.call("EXPIRE", KEYS[2], ARGV[5]) +end return 1 ` -func (c *billingCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { - key := userPlatformQuotaCacheKey(userID, platform) - _, err := c.rdb.Eval(ctx, updateUserPlatformQuotaUsageScript, []string{key}, +// userPlatformQuotaDirtySetKey 返回脏集(dirty set)的 Redis key。 +// 使用与 userPlatformQuotaCacheKey 相同的前缀 "billing:"。 +func userPlatformQuotaDirtySetKey() string { return "billing:" + "upq:dirty" } + +// userPlatformQuotaDirtyTTLSeconds 脏集兜底 TTL(秒):初始 SADD(Lua)与 Readd 共用, +// 确保 flusher 长期停摆时脏集最终过期;正常运行因持续 SADD 不断续期。 +const userPlatformQuotaDirtyTTLSeconds = 86400 + +// userPlatformQuotaDirtyMember 构造脏集成员字符串 "userID:platform"。 +func userPlatformQuotaDirtyMember(userID int64, platform string) string { + return strconv.FormatInt(userID, 10) + ":" + platform +} + +func (c *billingCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { + member := "" + if markDirty { + member = userPlatformQuotaDirtyMember(userID, platform) + } + _, err := c.rdb.Eval(ctx, updateUserPlatformQuotaUsageScript, + []string{userPlatformQuotaCacheKey(userID, platform), userPlatformQuotaDirtySetKey()}, strconv.FormatFloat(cost, 'f', -1, 64), int(ttl.Seconds()), service.UserPlatformQuotaCacheSchemaV1, + member, + userPlatformQuotaDirtyTTLSeconds, ).Result() if err != nil && !errors.Is(err, redis.Nil) { return err } return nil } + +// parseUserPlatformQuotaDirtyMember 将脏集成员字符串 "userID:platform" 解析为 +// service.UserPlatformQuotaKey。解析失败返回 ok=false。 +func parseUserPlatformQuotaDirtyMember(m string) (service.UserPlatformQuotaKey, bool) { + parts := strings.SplitN(m, ":", 2) + if len(parts) != 2 { + return service.UserPlatformQuotaKey{}, false + } + uid, err := strconv.ParseInt(parts[0], 10, 64) + if err != nil { + return service.UserPlatformQuotaKey{}, false + } + return service.UserPlatformQuotaKey{UserID: uid, Platform: parts[1]}, true +} + +// PopDirtyUserPlatformQuotaKeys 从脏集随机弹出最多 n 个 key。 +// 脏集为空时返回 (nil, nil)。 +func (c *billingCache) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]service.UserPlatformQuotaKey, error) { + members, err := c.rdb.SPopN(ctx, userPlatformQuotaDirtySetKey(), int64(n)).Result() + if err != nil { + if errors.Is(err, redis.Nil) { + return nil, nil + } + return nil, err + } + keys := make([]service.UserPlatformQuotaKey, 0, len(members)) + for _, m := range members { + k, ok := parseUserPlatformQuotaDirtyMember(m) + if !ok { + log.Printf("billing_cache: skipping invalid dirty member %q", m) + continue + } + keys = append(keys, k) + } + return keys, nil +} + +// ReaddDirtyUserPlatformQuotaKeys 将 keys 重新加入脏集(flush 失败时回填)。 +// 通过 pipeline 同时执行 SAdd + Expire,确保 Readd 后脏集具有兜底 TTL。 +// 空切片时直接返回 nil。 +func (c *billingCache) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []service.UserPlatformQuotaKey) error { + if len(keys) == 0 { + return nil + } + dirtyKey := userPlatformQuotaDirtySetKey() + members := make([]any, len(keys)) + for i, k := range keys { + members[i] = userPlatformQuotaDirtyMember(k.UserID, k.Platform) + } + pipe := c.rdb.Pipeline() + pipe.SAdd(ctx, dirtyKey, members...) + pipe.Expire(ctx, dirtyKey, userPlatformQuotaDirtyTTLSeconds*time.Second) + _, err := pipe.Exec(ctx) + return err +} + +// BatchGetUserPlatformQuotaCache 通过 Pipeline 批量 HGETALL 获取多个 user×platform 的 +// quota cache。返回切片与 keys 顺序、长度对齐;MISS 或解析失败位置返回 nil。 +func (c *billingCache) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []service.UserPlatformQuotaKey) ([]*service.UserPlatformQuotaCacheEntry, error) { + if len(keys) == 0 { + return nil, nil + } + pipe := c.rdb.Pipeline() + cmds := make([]*redis.MapStringStringCmd, len(keys)) + for i, k := range keys { + cmds[i] = pipe.HGetAll(ctx, userPlatformQuotaCacheKey(k.UserID, k.Platform)) + } + if _, err := pipe.Exec(ctx); err != nil && !errors.Is(err, redis.Nil) { + return nil, err + } + results := make([]*service.UserPlatformQuotaCacheEntry, len(keys)) + for i, cmd := range cmds { + m, err := cmd.Result() + if err != nil { + if !errors.Is(err, redis.Nil) { + log.Printf("billing_cache: BatchGet HGETALL cmd[%d] failed: %v (skip, self-heal)", i, err) + } + // 单个命令失败 → 对应位置 nil,继续 + continue + } + results[i] = parseUserPlatformQuotaHash(m) + } + return results, nil +} diff --git a/backend/internal/repository/billing_cache_user_platform_quota_test.go b/backend/internal/repository/billing_cache_user_platform_quota_test.go index 8d49fd31..15b185e7 100644 --- a/backend/internal/repository/billing_cache_user_platform_quota_test.go +++ b/backend/internal/repository/billing_cache_user_platform_quota_test.go @@ -88,7 +88,7 @@ func TestUserPlatformQuotaCache_NilLimitSetThenGet(t *testing.T) { func TestUserPlatformQuotaCache_IncrMissIsNoop(t *testing.T) { c, _ := newMiniRedisCache(t) - if err := c.IncrUserPlatformQuotaUsageCache(context.Background(), 1, "openai", 0.5, time.Minute); err != nil { + if err := c.IncrUserPlatformQuotaUsageCache(context.Background(), 1, "openai", 0.5, time.Minute, false); err != nil { t.Fatal(err) } _, ok, _ := c.GetUserPlatformQuotaCache(context.Background(), 1, "openai") @@ -105,10 +105,10 @@ func TestUserPlatformQuotaCache_IncrHitAccumulates(t *testing.T) { Version: 1, SchemaVersion: service.UserPlatformQuotaCacheSchemaV1, }, time.Minute) - if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.5, time.Minute); err != nil { + if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.5, time.Minute, false); err != nil { t.Fatal(err) } - if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.25, time.Minute); err != nil { + if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.25, time.Minute, false); err != nil { t.Fatal(err) } got, _, _ := c.GetUserPlatformQuotaCache(ctx, 1, "openai") diff --git a/backend/internal/repository/user_platform_quota_adapter_test.go b/backend/internal/repository/user_platform_quota_adapter_test.go index a55d2e9c..f31defe5 100644 --- a/backend/internal/repository/user_platform_quota_adapter_test.go +++ b/backend/internal/repository/user_platform_quota_adapter_test.go @@ -38,6 +38,9 @@ func (f *fakeRepoForAdapter) UpsertForUser(_ context.Context, userID int64, reco f.upsertCalledWith = records return f.upsertErr } +func (f *fakeRepoForAdapter) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error { + return nil +} func TestGenericAdapter_UpsertForUser_ForwardsRecords(t *testing.T) { fake := &fakeRepoForAdapter{} diff --git a/backend/internal/repository/user_platform_quota_repo.go b/backend/internal/repository/user_platform_quota_repo.go index 1e2e7f51..ccba2330 100644 --- a/backend/internal/repository/user_platform_quota_repo.go +++ b/backend/internal/repository/user_platform_quota_repo.go @@ -2,6 +2,7 @@ package repository import ( "context" + "errors" "fmt" "strings" "time" @@ -9,6 +10,7 @@ import ( dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/userplatformquota" "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/lib/pq" ) // UserPlatformQuotaRecord 是 repository 层的传输结构体, @@ -30,6 +32,22 @@ type UserPlatformQuotaRecord struct { // ErrUserPlatformQuotaNotFound 用于 ResetExpiredWindow 等需要"必须命中已有记录"的方法。 var ErrUserPlatformQuotaNotFound = fmt.Errorf("user platform quota record not found") +// ErrUserPlatformQuotaFKViolation 当批量 UPSERT 中存在 user_id 不在 users 表的记录时返回。 +var ErrUserPlatformQuotaFKViolation = errors.New("user platform quota snapshot FK violation") + +// UserPlatformQuotaSnapshot 是 BatchSnapshotUsage 的输入结构体, +// 表示 Redis 当前窗口快照(用于绝对值覆盖写入 DB)。 +type UserPlatformQuotaSnapshot struct { + UserID int64 + Platform string + DailyUsageUSD float64 + WeeklyUsageUSD float64 + MonthlyUsageUSD float64 + DailyWindowStart time.Time + WeeklyWindowStart time.Time + MonthlyWindowStart time.Time +} + // UserPlatformQuotaRepository 定义用户平台配额的数据访问接口。 type UserPlatformQuotaRepository interface { // BulkInsertInitial 幂等批量插入初始配额记录(ON CONFLICT DO NOTHING)。 @@ -44,6 +62,10 @@ type UserPlatformQuotaRepository interface { ResetExpiredWindow(ctx context.Context, userID int64, platform string, window string, newStart time.Time) error // UpsertForUser 全量替换该用户所有平台限额配置(详见 service.UserPlatformQuotaRepository.UpsertForUser)。 UpsertForUser(ctx context.Context, userID int64, records []UserPlatformQuotaRecord) error + // BatchSnapshotUsage 用一条多行 UPSERT 把整批 usage 以绝对值覆盖写入(非累加)。 + // usage/window_start 直接取 EXCLUDED(Redis 当前窗口快照),无 CASE。整批共用 now 作 created/updated_at。 + // 要求 snapshots 内 (user,platform) 不重复。FK 违反返回 ErrUserPlatformQuotaFKViolation。 + BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error } type userPlatformQuotaRepository struct { @@ -414,3 +436,73 @@ func insertLimitsRow(ctx context.Context, client *dbent.Client, userID int64, re } return nil } + +// batchRows 是 BatchSnapshotUsage 每批最大行数(9 参/行 × 6000 ≈ 54000 参,低于 Postgres 65535 上限)。 +const batchRows = 6000 + +// BatchSnapshotUsage 用一条多行 UPSERT 把整批 usage 以绝对值覆盖写入(非累加)。 +// 每批最多 batchRows 行;$1=now 共用;每行 8 个 per-row 参(user_id, platform, 3×usage, 3×window_start)。 +// FK 违反(user_id 不存在)返回 ErrUserPlatformQuotaFKViolation。 +// +// 注意:snapshots 超过 batchRows 会分多条 SQL 执行且【非单事务】——若某子批 FK 失败, +// 先前子批已写入无法回滚。调用方(flusher)应保证单次 batchSize ≤ batchRows +// (默认 flush_batch_size=1000 < 6000,安全)。 +// 另注:启用 flusher 后,本绝对值覆盖与 admin 直写 DB(ResetExpiredWindow/UpsertForUser)存在覆盖竞态, +// 详见 service/user_platform_quota_flusher.go 中 flushOneBatch 的"已知竞态"注释。 +func (r *userPlatformQuotaRepository) BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error { + if len(snapshots) == 0 { + return nil + } + + client := clientFromContext(ctx, r.client) + + for start := 0; start < len(snapshots); start += batchRows { + end := start + batchRows + if end > len(snapshots) { + end = len(snapshots) + } + batch := snapshots[start:end] + + var sb strings.Builder + _, _ = sb.WriteString( + "INSERT INTO user_platform_quotas" + + " (user_id, platform, daily_usage_usd, weekly_usage_usd, monthly_usage_usd," + + " daily_window_start, weekly_window_start, monthly_window_start, created_at, updated_at)" + + " VALUES ") + + // $1 = now(共用);每行 8 个 per-row 参,从 $2 起连续编号。 + args := []any{now} + for i, s := range batch { + if i > 0 { + _, _ = sb.WriteString(",") + } + b := len(args) // 当前 per-row 第一个参数的 0-based 索引,实际占位符 = b+1 + fmt.Fprintf(&sb, "($%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d,$1,$1)", + b+1, b+2, b+3, b+4, b+5, b+6, b+7, b+8) + args = append(args, + s.UserID, s.Platform, + s.DailyUsageUSD, s.WeeklyUsageUSD, s.MonthlyUsageUSD, + s.DailyWindowStart, s.WeeklyWindowStart, s.MonthlyWindowStart, + ) + } + + _, _ = sb.WriteString( + " ON CONFLICT (user_id, platform) WHERE deleted_at IS NULL DO UPDATE SET" + + " daily_usage_usd = EXCLUDED.daily_usage_usd," + + " weekly_usage_usd = EXCLUDED.weekly_usage_usd," + + " monthly_usage_usd = EXCLUDED.monthly_usage_usd," + + " daily_window_start = EXCLUDED.daily_window_start," + + " weekly_window_start = EXCLUDED.weekly_window_start," + + " monthly_window_start = EXCLUDED.monthly_window_start," + + " updated_at = EXCLUDED.updated_at") + + if _, err := client.ExecContext(ctx, sb.String(), args...); err != nil { + var pqErr *pq.Error + if errors.As(err, &pqErr) && pqErr.Code == "23503" { + return ErrUserPlatformQuotaFKViolation + } + return err + } + } + return nil +} diff --git a/backend/internal/repository/user_platform_quota_repo_integration_test.go b/backend/internal/repository/user_platform_quota_repo_integration_test.go index f02eeaa9..39e2f6e0 100644 --- a/backend/internal/repository/user_platform_quota_repo_integration_test.go +++ b/backend/internal/repository/user_platform_quota_repo_integration_test.go @@ -267,3 +267,101 @@ func TestUserPlatformQuotaRepository_ResetExpiredWindow_NotFoundReturnsSentinel( require.True(t, errors.Is(err, ErrUserPlatformQuotaNotFound), "expected ErrUserPlatformQuotaNotFound, got %v", err) } + +// TestBatchSnapshotUsage_InsertOverwriteMultiKey 验证 BatchSnapshotUsage 的绝对值覆盖语义: +// 1. 首批插入 2 条(不同 user),验证 daily 等于首批值; +// 2. 对同一 key 传不同值,验证 daily 等于新值(绝对覆盖,非累加)。 +func TestBatchSnapshotUsage_InsertOverwriteMultiKey(t *testing.T) { + ctx := context.Background() + // BatchSnapshotUsage 不开事务(直接写),使用独立 client 保证跨调用可见性。 + client := testEntClient(t) + + userID1 := mustCreateUserForQuota(t, client) + userID2 := mustCreateUserForQuota(t, client) + + repo := NewUserPlatformQuotaRepository(client) + + now := time.Date(2026, 5, 29, 12, 0, 0, 0, time.UTC) + dailyStart := time.Date(2026, 5, 29, 0, 0, 0, 0, time.UTC) + weeklyStart := time.Date(2026, 5, 25, 0, 0, 0, 0, time.UTC) // 当周一 + monthlyStart := time.Date(2026, 5, 1, 0, 0, 0, 0, time.UTC) + + // ── 第一批:插入 2 行 ────────────────────────────────────────────────────── + firstBatch := []UserPlatformQuotaSnapshot{ + { + UserID: userID1, + Platform: "anthropic", + DailyUsageUSD: 1.0, + WeeklyUsageUSD: 3.0, + MonthlyUsageUSD: 5.0, + DailyWindowStart: dailyStart, + WeeklyWindowStart: weeklyStart, + MonthlyWindowStart: monthlyStart, + }, + { + UserID: userID2, + Platform: "openai", + DailyUsageUSD: 2.0, + WeeklyUsageUSD: 4.0, + MonthlyUsageUSD: 6.0, + DailyWindowStart: dailyStart, + WeeklyWindowStart: weeklyStart, + MonthlyWindowStart: monthlyStart, + }, + } + require.NoError(t, repo.BatchSnapshotUsage(ctx, firstBatch, now), "first batch upsert") + + // 验证首批 daily 值 + rec1, err := repo.GetByUserPlatform(ctx, userID1, "anthropic") + require.NoError(t, err) + require.NotNil(t, rec1, "user1/anthropic should exist after first batch") + require.InDelta(t, 1.0, rec1.DailyUsageUSD, 1e-9, "user1 daily after first batch") + require.InDelta(t, 3.0, rec1.WeeklyUsageUSD, 1e-9, "user1 weekly after first batch") + require.InDelta(t, 5.0, rec1.MonthlyUsageUSD, 1e-9, "user1 monthly after first batch") + + rec2, err := repo.GetByUserPlatform(ctx, userID2, "openai") + require.NoError(t, err) + require.NotNil(t, rec2, "user2/openai should exist after first batch") + require.InDelta(t, 2.0, rec2.DailyUsageUSD, 1e-9, "user2 daily after first batch") + + // ── 第二批:对同一 key 传不同值,验证绝对覆盖(非累加)────────────────── + now2 := now.Add(5 * time.Minute) + secondBatch := []UserPlatformQuotaSnapshot{ + { + UserID: userID1, + Platform: "anthropic", + DailyUsageUSD: 9.9, // 新值,不是 1.0+9.9=10.9 + WeeklyUsageUSD: 19.9, // 新值,不是 3.0+19.9=22.9 + MonthlyUsageUSD: 29.9, // 新值 + DailyWindowStart: dailyStart, + WeeklyWindowStart: weeklyStart, + MonthlyWindowStart: monthlyStart, + }, + { + UserID: userID2, + Platform: "openai", + DailyUsageUSD: 8.8, + WeeklyUsageUSD: 18.8, + MonthlyUsageUSD: 28.8, + DailyWindowStart: dailyStart, + WeeklyWindowStart: weeklyStart, + MonthlyWindowStart: monthlyStart, + }, + } + require.NoError(t, repo.BatchSnapshotUsage(ctx, secondBatch, now2), "second batch upsert") + + // 验证第二批覆盖:daily 应为新值,不是累加 + rec1After, err := repo.GetByUserPlatform(ctx, userID1, "anthropic") + require.NoError(t, err) + require.NotNil(t, rec1After) + require.InDelta(t, 9.9, rec1After.DailyUsageUSD, 1e-9, "user1 daily must be overwritten to 9.9 (not accumulated)") + require.InDelta(t, 19.9, rec1After.WeeklyUsageUSD, 1e-9, "user1 weekly must be overwritten to 19.9") + require.InDelta(t, 29.9, rec1After.MonthlyUsageUSD, 1e-9, "user1 monthly must be overwritten to 29.9") + + rec2After, err := repo.GetByUserPlatform(ctx, userID2, "openai") + require.NoError(t, err) + require.NotNil(t, rec2After) + require.InDelta(t, 8.8, rec2After.DailyUsageUSD, 1e-9, "user2 daily must be overwritten to 8.8 (not accumulated)") + require.InDelta(t, 18.8, rec2After.WeeklyUsageUSD, 1e-9, "user2 weekly must be overwritten to 18.8") + require.InDelta(t, 28.8, rec2After.MonthlyUsageUSD, 1e-9, "user2 monthly must be overwritten to 28.8") +} diff --git a/backend/internal/repository/user_platform_quota_service_adapter.go b/backend/internal/repository/user_platform_quota_service_adapter.go index 7495cd26..5240bb54 100644 --- a/backend/internal/repository/user_platform_quota_service_adapter.go +++ b/backend/internal/repository/user_platform_quota_service_adapter.go @@ -94,6 +94,29 @@ func (a *userPlatformQuotaServiceAdapter) ResetExpiredWindow(ctx context.Context return err } +// BatchSnapshotUsage 转换 []service.UserPlatformQuotaSnapshot → []UserPlatformQuotaSnapshot, +// 调底层 repo,并将 repository FK sentinel 包装为 service sentinel。 +func (a *userPlatformQuotaServiceAdapter) BatchSnapshotUsage(ctx context.Context, snapshots []service.UserPlatformQuotaSnapshot, now time.Time) error { + repoSnaps := make([]UserPlatformQuotaSnapshot, len(snapshots)) + for i, s := range snapshots { + repoSnaps[i] = UserPlatformQuotaSnapshot{ + UserID: s.UserID, + Platform: s.Platform, + DailyUsageUSD: s.DailyUsageUSD, + WeeklyUsageUSD: s.WeeklyUsageUSD, + MonthlyUsageUSD: s.MonthlyUsageUSD, + DailyWindowStart: s.DailyWindowStart, + WeeklyWindowStart: s.WeeklyWindowStart, + MonthlyWindowStart: s.MonthlyWindowStart, + } + } + err := a.inner.BatchSnapshotUsage(ctx, repoSnaps, now) + if errors.Is(err, ErrUserPlatformQuotaFKViolation) { + return fmt.Errorf("%w: %v", service.ErrUserPlatformQuotaFKViolation, err) + } + return err +} + // genericUserPlatformQuotaAdapter 通过通用接口适配(用于测试 fake 或非标准实现)。 type genericUserPlatformQuotaAdapter struct { inner UserPlatformQuotaRepository @@ -167,6 +190,29 @@ func (a *genericUserPlatformQuotaAdapter) ResetExpiredWindow(ctx context.Context return err } +// BatchSnapshotUsage 转换 []service.UserPlatformQuotaSnapshot → []UserPlatformQuotaSnapshot(通用 adapter), +// 并将 repository FK sentinel 包装为 service sentinel。 +func (a *genericUserPlatformQuotaAdapter) BatchSnapshotUsage(ctx context.Context, snapshots []service.UserPlatformQuotaSnapshot, now time.Time) error { + repoSnaps := make([]UserPlatformQuotaSnapshot, len(snapshots)) + for i, s := range snapshots { + repoSnaps[i] = UserPlatformQuotaSnapshot{ + UserID: s.UserID, + Platform: s.Platform, + DailyUsageUSD: s.DailyUsageUSD, + WeeklyUsageUSD: s.WeeklyUsageUSD, + MonthlyUsageUSD: s.MonthlyUsageUSD, + DailyWindowStart: s.DailyWindowStart, + WeeklyWindowStart: s.WeeklyWindowStart, + MonthlyWindowStart: s.MonthlyWindowStart, + } + } + err := a.inner.BatchSnapshotUsage(ctx, repoSnaps, now) + if errors.Is(err, ErrUserPlatformQuotaFKViolation) { + return fmt.Errorf("%w: %v", service.ErrUserPlatformQuotaFKViolation, err) + } + return err +} + // toServiceRecord 将 repository.UserPlatformQuotaRecord 转换为 service.UserPlatformQuotaRecord。 func toServiceRecord(rec *UserPlatformQuotaRecord) *service.UserPlatformQuotaRecord { return &service.UserPlatformQuotaRecord{ diff --git a/backend/internal/service/admin_service_delete_test.go b/backend/internal/service/admin_service_delete_test.go index d01b11e6..2aae73a9 100644 --- a/backend/internal/service/admin_service_delete_test.go +++ b/backend/internal/service/admin_service_delete_test.go @@ -471,10 +471,22 @@ func (s *billingCacheStub) DeleteUserPlatformQuotaCache(ctx context.Context, use panic("unexpected DeleteUserPlatformQuotaCache call") } -func (s *billingCacheStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { +func (s *billingCacheStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { panic("unexpected IncrUserPlatformQuotaUsageCache call") } +func (s *billingCacheStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) { + panic("unexpected PopDirtyUserPlatformQuotaKeys call") +} + +func (s *billingCacheStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error { + panic("unexpected ReaddDirtyUserPlatformQuotaKeys call") +} + +func (s *billingCacheStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + panic("unexpected BatchGetUserPlatformQuotaCache call") +} + func waitForInvalidations(t *testing.T, ch <-chan subscriptionInvalidateCall, expected int) []subscriptionInvalidateCall { t.Helper() calls := make([]subscriptionInvalidateCall, 0, expected) diff --git a/backend/internal/service/auth_service_platform_quota_test.go b/backend/internal/service/auth_service_platform_quota_test.go index f58dc48c..46069814 100644 --- a/backend/internal/service/auth_service_platform_quota_test.go +++ b/backend/internal/service/auth_service_platform_quota_test.go @@ -43,6 +43,10 @@ func (f *fakeInsertRecorder) ResetExpiredWindow(_ context.Context, _ int64, _ st return nil } +func (f *fakeInsertRecorder) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error { + return nil +} + func TestSnapshotPlatformQuotaDefaults_PassesToRepoBulkInsert(t *testing.T) { fakeRepo := &fakeInsertRecorder{} s := &AuthService{userPlatformQuotaRepo: fakeRepo} diff --git a/backend/internal/service/auth_service_register_test.go b/backend/internal/service/auth_service_register_test.go index a7c0d260..2ee9f21a 100644 --- a/backend/internal/service/auth_service_register_test.go +++ b/backend/internal/service/auth_service_register_test.go @@ -105,6 +105,10 @@ func (s *userPlatformQuotaRepoStub) ResetExpiredWindow(context.Context, int64, s panic("unexpected ResetExpiredWindow call") } +func (s *userPlatformQuotaRepoStub) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error { + return nil +} + func (s *defaultSubscriptionAssignerStub) AssignOrExtendSubscription(_ context.Context, input *AssignSubscriptionInput) (*UserSubscription, bool, error) { if input != nil { s.calls = append(s.calls, *input) diff --git a/backend/internal/service/billing_cache_service.go b/backend/internal/service/billing_cache_service.go index 8a5172f4..b734fab1 100644 --- a/backend/internal/service/billing_cache_service.go +++ b/backend/internal/service/billing_cache_service.go @@ -689,7 +689,8 @@ func (s *BillingCacheService) IncrementUserPlatformQuotaUsage(userID int64, plat ctx, cancel := context.WithTimeout(context.Background(), cacheWriteTimeout) defer cancel() ttl := time.Duration(s.cfg.Billing.UserPlatformQuotaCacheTTLSeconds) * time.Second - if err := s.cache.IncrUserPlatformQuotaUsageCache(ctx, userID, platform, cost, ttl); err != nil { + markDirty := s.cfg.Database.UserPlatformQuotaFlusherEnabled + if err := s.cache.IncrUserPlatformQuotaUsageCache(ctx, userID, platform, cost, ttl, markDirty); err != nil { logger.LegacyPrintf("service.billing_cache", "ALERT: incr user platform quota cache failed user=%d platform=%s cost=%f: %v", userID, platform, cost, err) @@ -1310,3 +1311,20 @@ func monthlyQuotaWindowExpired(start *time.Time, now time.Time) bool { } return now.Sub(*start) >= 30*24*time.Hour } + +// HasUserPlatformQuotaLimit 判断该 user×platform 是否设了任一非 nil limit。 +// 写入点守卫:无 limit 直接跳过 Redis 写 + 脏集标记,消除无谓写入。 +// fail-safe:任何不确定(simple 模式除外)都返回 true 维持写入。 +func (s *BillingCacheService) HasUserPlatformQuotaLimit(ctx context.Context, userID int64, platform string) bool { + if s.cfg.RunMode == config.RunModeSimple { + return false + } + if s.cache == nil { + return true + } + entry, ok, err := s.cache.GetUserPlatformQuotaCache(ctx, userID, platform) + if err != nil || !ok || entry == nil { + return true + } + return entry.DailyLimitUSD != nil || entry.WeeklyLimitUSD != nil || entry.MonthlyLimitUSD != nil +} diff --git a/backend/internal/service/billing_cache_service_singleflight_test.go b/backend/internal/service/billing_cache_service_singleflight_test.go index b443d97e..235b13a6 100644 --- a/backend/internal/service/billing_cache_service_singleflight_test.go +++ b/backend/internal/service/billing_cache_service_singleflight_test.go @@ -79,10 +79,22 @@ func (s *billingCacheMissStub) DeleteUserPlatformQuotaCache(ctx context.Context, return nil } -func (s *billingCacheMissStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { +func (s *billingCacheMissStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { return nil } +func (s *billingCacheMissStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) { + return nil, nil +} + +func (s *billingCacheMissStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error { + return nil +} + +func (s *billingCacheMissStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + return nil, nil +} + type balanceLoadUserRepoStub struct { mockUserRepo calls atomic.Int64 diff --git a/backend/internal/service/billing_cache_service_test.go b/backend/internal/service/billing_cache_service_test.go index bcd086fa..c344b417 100644 --- a/backend/internal/service/billing_cache_service_test.go +++ b/backend/internal/service/billing_cache_service_test.go @@ -80,10 +80,22 @@ func (b *billingCacheWorkerStub) DeleteUserPlatformQuotaCache(ctx context.Contex return nil } -func (b *billingCacheWorkerStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { +func (b *billingCacheWorkerStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { return nil } +func (b *billingCacheWorkerStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) { + return nil, nil +} + +func (b *billingCacheWorkerStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error { + return nil +} + +func (b *billingCacheWorkerStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + return nil, nil +} + func TestBillingCacheServiceQueueHighLoad(t *testing.T) { cache := &billingCacheWorkerStub{} svc := NewBillingCacheService(cache, nil, nil, nil, nil, nil, &config.Config{}, nil) diff --git a/backend/internal/service/billing_cache_service_user_platform_quota_test.go b/backend/internal/service/billing_cache_service_user_platform_quota_test.go index 674aa9a5..a82c0e00 100644 --- a/backend/internal/service/billing_cache_service_user_platform_quota_test.go +++ b/backend/internal/service/billing_cache_service_user_platform_quota_test.go @@ -20,14 +20,15 @@ type fakeIncrCache struct { } type incrCall struct { - userID int64 - platform string - cost float64 - ttl time.Duration + userID int64 + platform string + cost float64 + ttl time.Duration + markDirty bool } -func (f *fakeIncrCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { - f.calls = append(f.calls, incrCall{userID, platform, cost, ttl}) +func (f *fakeIncrCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { + f.calls = append(f.calls, incrCall{userID, platform, cost, ttl, markDirty}) return nil } @@ -49,10 +50,10 @@ func TestIncrementUserPlatformQuotaUsage_SyncCallsCache(t *testing.T) { if len(fake.calls) != 2 { t.Fatalf("expected 2 incr calls, got %d", len(fake.calls)) } - if fake.calls[0] != (incrCall{101, "anthropic", 0.25, 120 * time.Second}) { + if fake.calls[0] != (incrCall{userID: 101, platform: "anthropic", cost: 0.25, ttl: 120 * time.Second, markDirty: false}) { t.Errorf("call[0] = %+v", fake.calls[0]) } - if fake.calls[1] != (incrCall{101, "openai", 0.50, 120 * time.Second}) { + if fake.calls[1] != (incrCall{userID: 101, platform: "openai", cost: 0.50, ttl: 120 * time.Second, markDirty: false}) { t.Errorf("call[1] = %+v", fake.calls[1]) } } @@ -88,7 +89,11 @@ func (f *fakeQuotaRepo) ResetExpiredWindow(_ context.Context, _ int64, _ string, return nil } -// fakeFullCache 同时支持 Get + Set + Incr + Delete。 +func (f *fakeQuotaRepo) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error { + return nil +} + +// fakeFullCache 同时支持 Get + Set + Incr + Delete + Pop/Readd/BatchGet(脏集读写)。 // mu 保护 entry 和 deleteCalls,防止异步 goroutine 与主 goroutine 之间的 data race。 type fakeFullCache struct { BillingCache @@ -99,6 +104,8 @@ type fakeFullCache struct { lastSetTTL time.Duration // 最近一次 Set 的 ttl getErr error // 非 nil 时 Get 先返回 (nil,false,getErr) setErr error // 非 nil 时 Set 返回该 err(setCalls 仍+1) + // dirty 模拟脏集,供 flusher 测试使用。 + dirty map[UserPlatformQuotaKey]struct{} } // getDeleteCalls 线程安全地读取 deleteCalls。 @@ -161,6 +168,48 @@ func (f *fakeFullCache) DeleteUserPlatformQuotaCache(_ context.Context, _ int64, return nil } +func (f *fakeFullCache) PopDirtyUserPlatformQuotaKeys(_ context.Context, n int) ([]UserPlatformQuotaKey, error) { + f.mu.Lock() + defer f.mu.Unlock() + if len(f.dirty) == 0 { + return nil, nil + } + keys := make([]UserPlatformQuotaKey, 0, n) + for k := range f.dirty { + if len(keys) >= n { + break + } + keys = append(keys, k) + delete(f.dirty, k) + } + return keys, nil +} + +func (f *fakeFullCache) ReaddDirtyUserPlatformQuotaKeys(_ context.Context, keys []UserPlatformQuotaKey) error { + f.mu.Lock() + defer f.mu.Unlock() + if f.dirty == nil { + f.dirty = make(map[UserPlatformQuotaKey]struct{}) + } + for _, k := range keys { + f.dirty[k] = struct{}{} + } + return nil +} + +// BatchGetUserPlatformQuotaCache 对每个 key 返回 f.entry(MISS → nil), +// 保持与输入 keys 顺序/长度对齐。注意此处所有 key 共享同一个 entry, +// 仅用于测试场景。 +func (f *fakeFullCache) BatchGetUserPlatformQuotaCache(_ context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + f.mu.Lock() + defer f.mu.Unlock() + results := make([]*UserPlatformQuotaCacheEntry, len(keys)) + for i := range keys { + results[i] = f.entry + } + return results, nil +} + func newServiceForPreflight(t *testing.T, repo UserPlatformQuotaRepository, cache BillingCache) *BillingCacheService { t.Helper() cfg := &config.Config{} @@ -713,3 +762,71 @@ func TestCheckUserPlatformQuotaEligibility_SentinelCrossDay_NoRefresh(t *testing t.Errorf("sentinel cross-window must NOT trigger refresh SetCache, got %d calls", cache.getSetCalls()) } } + +// ── TestHasUserPlatformQuotaLimit ──────────────────────────────────────────── + +func TestHasUserPlatformQuotaLimit(t *testing.T) { + daily := 5.0 + + tests := []struct { + name string + setup func() *BillingCacheService + want bool + }{ + { + name: "has_limit", + setup: func() *BillingCacheService { + entry := &UserPlatformQuotaCacheEntry{DailyLimitUSD: &daily} + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry}) + return svc + }, + want: true, + }, + { + name: "sentinel_no_limit", + setup: func() *BillingCacheService { + entry := &UserPlatformQuotaCacheEntry{} // 三个 limit 字段全 nil + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry}) + return svc + }, + want: false, + }, + { + name: "cache_miss", + setup: func() *BillingCacheService { + // entry==nil → GetUserPlatformQuotaCache 返回 (nil,false,nil) + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{}) + return svc + }, + want: true, // fail-safe + }, + { + name: "redis_err", + setup: func() *BillingCacheService { + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{getErr: errors.New("redis down")}) + return svc + }, + want: true, // fail-safe + }, + { + name: "simple_mode", + setup: func() *BillingCacheService { + entry := &UserPlatformQuotaCacheEntry{DailyLimitUSD: &daily} + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry}) + svc.cfg.RunMode = config.RunModeSimple + return svc + }, + want: false, // simple 模式始终跳过 + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := tt.setup() + got := svc.HasUserPlatformQuotaLimit(context.Background(), 1, "anthropic") + if got != tt.want { + t.Errorf("HasUserPlatformQuotaLimit() = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 940a827d..6b1438e8 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -21,6 +21,12 @@ type APIKeyRateLimitCacheData struct { Window7d int64 `json:"window_7d"` } +// UserPlatformQuotaKey 标识一个 user×platform,用于脏集出入与批量读。 +type UserPlatformQuotaKey struct { + UserID int64 + Platform string +} + // UserPlatformQuotaCacheEntry Redis hash 反序列化结果。 // // SchemaVersion 用于向后兼容: @@ -72,7 +78,13 @@ type BillingCache interface { SetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string, entry *UserPlatformQuotaCacheEntry, ttl time.Duration) error DeleteUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) error // IncrUserPlatformQuotaUsageCache 在缓存命中时累加用量;缓存未命中(key 不存在)静默返回 nil。 - IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error + // markDirty=true 时将该 key 的 member 写入 Redis 脏集,供 flusher 批量回写 DB。 + IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error + + // 脏集读写,供 flusher 使用。 + PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) + ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error + BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) } // ModelPricing 模型价格配置(per-token价格,与LiteLLM格式一致) diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index effa803a..f807f3ec 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -96,11 +96,13 @@ var ( modelsListCacheMissTotal atomic.Int64 modelsListCacheStoreTotal atomic.Int64 + // Deprecated: flusher_enabled=true 后不再增长(仅 flag=false 降级直写路径使用);新主路径见 FlusherMetrics。remove after 2026-09。 // userPlatformQuotaDBIncrErrorTotal 统计 finalizePostUsageBilling 异步 goroutine // 中 IncrementUsageWithReset 失败次数。Redis 已成功累加 + DB 写失败意味着 // Redis cache TTL 过期或被清后该笔 cost 会丢失(与实际消费偏差)。 // oncall 通过 GatewayUserPlatformQuotaIncrStats() 暴露给 ops 面板做阈值告警。 userPlatformQuotaDBIncrErrorTotal atomic.Int64 + // Deprecated: flusher_enabled=true 后不再增长(仅 flag=false 降级直写路径使用);新主路径见 FlusherMetrics。remove after 2026-09。 // userPlatformQuotaDBIncrLegacyErrorTotal 统计 legacy postUsageBilling // (applyUsageBilling 在 repo==nil 时 fallback)路径下的失败次数; // 与 DB Incr 失败分开计数,便于区分"主路径暂时故障"vs"基础设施长期未配齐"。 @@ -141,6 +143,23 @@ func GatewayUserPlatformQuotaIncrStats() (mainPathErr, legacyPathErr, sentinelSe userPlatformQuotaSentinelSetCacheErrorTotal.Load() } +// GatewayUserPlatformQuotaFlusherStats 暴露 flusher 运行指标供 ops/health 面板查询。 +func GatewayUserPlatformQuotaFlusherStats(f *UserPlatformQuotaUsageFlusher) map[string]int64 { + if f == nil || f.metrics == nil { + return nil + } + m := f.metrics + return map[string]int64{ + "flush_success": m.FlushSuccessTotal.Load(), + "flush_error": m.FlushErrorTotal.Load(), + "flush_batch_size": m.FlushBatchSizeTotal.Load(), + "flush_latency_ms_max": m.FlushLatencyMsMax.Load(), + "dirty_readd": m.DirtyReaddTotal.Load(), + "dirty_lost": m.DirtyLostTotal.Load(), + "flush_fk_violation": m.FlushFKViolationTotal.Load(), + } +} + func openAIStreamEventIsTerminal(data string) bool { trimmed := strings.TrimSpace(data) if trimmed == "" { @@ -8234,18 +8253,23 @@ func postUsageBilling(ctx context.Context, p *postUsageBillingParams, deps *bill } } - // Platform quota DB-only 累加(与 finalizePostUsageBilling 行为对齐的兜底): - // - 仅对 standard(余额)模式生效;订阅模式豁免 - // - 直接走 DB,不经 Redis Incr 队列:legacy 路径在 repo==nil(仓库未注入) - // 时被触发,此时整套 billing repo 都不可用,没有"双队列"风险 - // - 失败仅记 ALERT log + counter,不阻断主扣费流程;与正常路径一致 - // - // 历史背景:原 legacy path 完全跳过此累加,导致部署中如果 repo 偶然为 nil - // 时用户消费可绕过 platform quota,存在静默资金风险。 + // Platform quota 累加(legacy 兜底路径):仅对 standard(余额)模式生效;订阅模式豁免;仅对有 limit 的用户写 + // - HasUserPlatformQuotaLimit 守卫:与正常路径对齐,无 limit 公司跳过 + // - 新增 Redis 同步写:enforcement 走 Redis,legacy 路径也必须同步写,否则 preflight 看不到消费 + // - flusher_enabled=false(降级):保留原有同步直写 DB + // - flusher_enabled=true:跳过直写 DB,由 flusher 异步批量刷(markDirty 在 IncrementUserPlatformQuotaUsage 内部完成) + // - 失败仅记 ALERT log + counter,不阻断主扣费流程 if !p.IsSubscriptionBill && p.Platform != "" && cost.ActualCost > 0 && p.User != nil && deps.userPlatformQuotaRepo != nil { - if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(billingCtx, p.User.ID, p.Platform, cost.ActualCost, time.Now().UTC()); err != nil { - userPlatformQuotaDBIncrLegacyErrorTotal.Add(1) - logger.LegacyPrintf("service.gateway", "ALERT: legacy incr user platform quota DB failed user=%d platform=%s cost=%f: %v", p.User.ID, p.Platform, cost.ActualCost, err) + if deps.billingCacheService.HasUserPlatformQuotaLimit(billingCtx, p.User.ID, p.Platform) { + deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, cost.ActualCost) + if deps.cfg == nil || !deps.cfg.Database.UserPlatformQuotaFlusherEnabled { + // 降级路径:flusher 未启用时保留原有同步直写 DB + if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(billingCtx, p.User.ID, p.Platform, cost.ActualCost, time.Now().UTC()); err != nil { + userPlatformQuotaDBIncrLegacyErrorTotal.Add(1) + logger.LegacyPrintf("service.gateway", "ALERT: legacy incr user platform quota DB failed user=%d platform=%s cost=%f: %v", p.User.ID, p.Platform, cost.ActualCost, err) + } + } + // flusher_enabled=true:不直写 DB,flusher 异步批量刷 } } @@ -8395,30 +8419,38 @@ func finalizePostUsageBilling(ctx context.Context, p *postUsageBillingParams, de deps.deferredService.ScheduleLastUsedUpdate(p.Account.ID) - // Platform quota 累加:仅在 standard(余额)模式生效;订阅模式豁免 - // Redis 同步写 + DB 异步持久化: + // Platform quota 累加:仅在 standard(余额)模式生效;订阅模式豁免;仅对有 limit 的用户写 + // Redis 同步写 + DB 异步持久化(flag=false 降级)或 flusher 异步刷(flag=true): + // - HasUserPlatformQuotaLimit 守卫:无 limit 的公司跳过,避免无效写入 + 浪费 Redis 容量 // - Redis 同步:确保下次 preflight 立即看到最新 usage,把 TOCTOU 超支窗口 // 限制在并发 in-flight 请求数量内(旧实现的异步入队会让超支无限累积直到 worker 处理) - // - DB 异步:在独立 goroutine 中走 detached context,失败用 ALERT log 触发 oncall 对账 + // - DB 异步(flusher_enabled=false):在独立 goroutine 中走 detached context,失败用 ALERT log 触发 oncall 对账 + // - flusher_enabled=true:不直写 DB,由 flusher 异步批量刷(markDirty 已在 IncrementUserPlatformQuotaUsage 内部完成) if !p.IsSubscriptionBill && p.Platform != "" && p.Cost.ActualCost > 0 && p.User != nil && deps.userPlatformQuotaRepo != nil { - deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, p.Cost.ActualCost) - dbCtx, dbCancel := detachUpstreamContext(ctx) - userID, platform, cost := p.User.ID, p.Platform, p.Cost.ActualCost - go func() { - defer func() { - if r := recover(); r != nil { - logger.LegacyPrintf("service.gateway", "ALERT: panic in user platform quota incr goroutine user=%d platform=%s: %v", userID, platform, r) - } - }() - defer dbCancel() - if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(dbCtx, userID, platform, cost, time.Now().UTC()); err != nil { - // 失败计数器:暴露给 GatewayUserPlatformQuotaIncrStats(),由 ops 面板做斜率告警。 - userPlatformQuotaDBIncrErrorTotal.Add(1) - // ALERT 级别:DB 持久化失败意味着 Redis cache 失效后该笔 cost 永久丢失, - // 用户配额视图与实际消费会偏差,oncall 需要据此对账或人工补录。 - logger.LegacyPrintf("service.gateway", "ALERT: incr user platform quota DB failed user=%d platform=%s cost=%f: %v", userID, platform, cost, err) + if deps.billingCacheService.HasUserPlatformQuotaLimit(ctx, p.User.ID, p.Platform) { + deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, p.Cost.ActualCost) + if deps.cfg == nil || !deps.cfg.Database.UserPlatformQuotaFlusherEnabled { + // 降级路径:flusher 未启用时保留原有异步直写 DB + dbCtx, dbCancel := detachUpstreamContext(ctx) + userID, platform, cost := p.User.ID, p.Platform, p.Cost.ActualCost + go func() { + defer func() { + if r := recover(); r != nil { + logger.LegacyPrintf("service.gateway", "ALERT: panic in user platform quota incr goroutine user=%d platform=%s: %v", userID, platform, r) + } + }() + defer dbCancel() + if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(dbCtx, userID, platform, cost, time.Now().UTC()); err != nil { + // 失败计数器:暴露给 GatewayUserPlatformQuotaIncrStats(),由 ops 面板做斜率告警。 + userPlatformQuotaDBIncrErrorTotal.Add(1) + // ALERT 级别:DB 持久化失败意味着 Redis cache 失效后该笔 cost 永久丢失, + // 用户配额视图与实际消费会偏差,oncall 需要据此对账或人工补录。 + logger.LegacyPrintf("service.gateway", "ALERT: incr user platform quota DB failed user=%d platform=%s cost=%f: %v", userID, platform, cost, err) + } + }() } - }() + // flusher_enabled=true:不直写 DB,flusher 异步批量刷 + } } // Notification checks run async — all parameters are already captured, @@ -8533,6 +8565,7 @@ type billingDeps struct { deferredService *DeferredService balanceNotifyService *BalanceNotifyService userPlatformQuotaRepo UserPlatformQuotaRepository + cfg *config.Config } func (s *GatewayService) billingDeps() *billingDeps { @@ -8544,6 +8577,7 @@ func (s *GatewayService) billingDeps() *billingDeps { deferredService: s.deferredService, balanceNotifyService: s.balanceNotifyService, userPlatformQuotaRepo: s.userPlatformQuotaRepo, + cfg: s.cfg, } } diff --git a/backend/internal/service/user_platform_quota_flusher.go b/backend/internal/service/user_platform_quota_flusher.go new file mode 100644 index 00000000..3ee23d2c --- /dev/null +++ b/backend/internal/service/user_platform_quota_flusher.go @@ -0,0 +1,267 @@ +package service + +import ( + "context" + "errors" + "sync/atomic" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +// quotaDirtyCache 是 flusher 依赖的窄接口(来自 BillingCache)。 +type quotaDirtyCache interface { + PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) + ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error + BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) +} + +// quotaSnapshotWriter 是 flusher 依赖的 DB 写入窄接口。 +// 使用 service 层的 UserPlatformQuotaSnapshot,避免与 repository 包形成循环依赖; +// 实际实现由 repository adapter 在 B7 注入。 +type quotaSnapshotWriter interface { + BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error +} + +// FlusherMetrics 记录 flusher 运行时指标(原子量,零值可用)。 +type FlusherMetrics struct { + FlushSuccessTotal atomic.Int64 + FlushErrorTotal atomic.Int64 + FlushBatchSizeTotal atomic.Int64 + FlushLatencyMsMax atomic.Int64 + DirtyReaddTotal atomic.Int64 + // DirtyLostTotal:Readd 失败导致脏 key 丢失——已 SPOP+主操作失败+Readd 也失败; + // Redis 仍权威,活跃 key 下次 SADD 自愈。 + DirtyLostTotal atomic.Int64 + FlushFKViolationTotal atomic.Int64 +} + +// flusherMaxBatchesPerTick 单次 tick 最多消费的批数,防止 tick 执行时间过长。 +const flusherMaxBatchesPerTick = 16 + +// maxFlushBatchSize 限制单批行数,必须 ≤ repository.BatchSnapshotUsage 的 batchRows(6000), +// 以保证单次 flush 的 snapshots 仅生成一条 UPSERT(单事务原子)。两处需手动保持一致。 +const maxFlushBatchSize = 6000 + +// defaultFlushBatchSize 是配置 flush_batch_size 非法(≤0)时的回退值。 +const defaultFlushBatchSize = 1000 + +// UserPlatformQuotaUsageFlusher 将 Redis 脏集快照定期批量写入 DB。 +// 不维护任何 delta/in-process 状态;每批读取 Redis 当前绝对值覆盖写入。 +type UserPlatformQuotaUsageFlusher struct { + cache quotaDirtyCache + quotaRepo quotaSnapshotWriter + timingWheel *TimingWheelService + // enabled 对应 flusher_enabled 配置;false 时 Start() 不注册定时器。 + enabled bool + interval time.Duration + batchSize int + flushTimeout time.Duration + metrics *FlusherMetrics + stopped atomic.Bool +} + +// NewUserPlatformQuotaUsageFlusher 创建 UserPlatformQuotaUsageFlusher。 +// cache(BillingCache) 隐式满足 quotaDirtyCache;quotaRepo(UserPlatformQuotaRepository) 隐式满足 quotaSnapshotWriter。 +func NewUserPlatformQuotaUsageFlusher(cfg *config.Config, cache BillingCache, quotaRepo UserPlatformQuotaRepository, tw *TimingWheelService) *UserPlatformQuotaUsageFlusher { + batchSize := cfg.Database.UserPlatformQuotaFlushBatchSize + if batchSize <= 0 { + batchSize = defaultFlushBatchSize + } + if batchSize > maxFlushBatchSize { + logger.LegacyPrintf("quota_flusher", + "[QuotaFlusher] flush_batch_size %d 超过上限 %d,已 clamp(避免 BatchSnapshotUsage 多子批非原子)", + cfg.Database.UserPlatformQuotaFlushBatchSize, maxFlushBatchSize) + batchSize = maxFlushBatchSize + } + interval := time.Duration(cfg.Database.UserPlatformQuotaFlushIntervalMs) * time.Millisecond + if interval <= 0 { + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] flush_interval_ms %d 非法,回退 2000ms", cfg.Database.UserPlatformQuotaFlushIntervalMs) + interval = 2 * time.Second + } + return &UserPlatformQuotaUsageFlusher{ + cache: cache, + quotaRepo: quotaRepo, + timingWheel: tw, + enabled: cfg.Database.UserPlatformQuotaFlusherEnabled, + interval: interval, + batchSize: batchSize, + flushTimeout: 3 * time.Second, + metrics: &FlusherMetrics{}, + } +} + +// updateLatencyMax 用 CAS 单调更新最大延迟。 +func (s *UserPlatformQuotaUsageFlusher) updateLatencyMax(ms int64) { + for { + old := s.metrics.FlushLatencyMsMax.Load() + if ms <= old { + return + } + if s.metrics.FlushLatencyMsMax.CompareAndSwap(old, ms) { + return + } + } +} + +// readdOrCountLost 尝试把 keys 回填脏集:成功计 DirtyReaddTotal,失败计 DirtyLostTotal 并 ALERT。 +func (s *UserPlatformQuotaUsageFlusher) readdOrCountLost(ctx context.Context, keys []UserPlatformQuotaKey, stage string) { + if err := s.cache.ReaddDirtyUserPlatformQuotaKeys(ctx, keys); err != nil { + s.metrics.DirtyLostTotal.Add(int64(len(keys))) + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] ALERT: Readd after %s failed, %d keys 丢出脏集(DB 镜像缺这批,Redis 仍权威,活跃 key 下次 SADD 自愈): %v", stage, len(keys), err) + return + } + s.metrics.DirtyReaddTotal.Add(int64(len(keys))) +} + +// flushOneBatch 处理单批:Pop → BatchGet → 组装 snaps → BatchSnapshotUsage。 +// 返回 (shouldContinue bool):false 表示本轮循环应停止(空集/错误/最后一批)。 +// 每次调用独立创建带 timeout 的 ctx 并 defer cancel,不会在循环中累积泄漏。 +func (s *UserPlatformQuotaUsageFlusher) flushOneBatch(parentCtx context.Context) bool { + ctx, cancel := context.WithTimeout(parentCtx, s.flushTimeout) + defer cancel() + + // 1. Pop 脏集 + keys, err := s.cache.PopDirtyUserPlatformQuotaKeys(ctx, s.batchSize) + if err != nil { + s.metrics.FlushErrorTotal.Add(1) + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] PopDirty error: %v", err) + return false + } + if len(keys) == 0 { + // 脏集已空 + return false + } + + // 2. 批量读 Redis 快照 + entries, err := s.cache.BatchGetUserPlatformQuotaCache(ctx, keys) + if err != nil { + s.metrics.FlushErrorTotal.Add(1) + s.readdOrCountLost(ctx, keys, "BatchGet") + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] BatchGet error: %v", err) + return false + } + + // 3. 组装 snapshots(MISS 或任一 WindowStart==nil → 跳过) + snaps := make([]UserPlatformQuotaSnapshot, 0, len(keys)) + for i, key := range keys { + e := entries[i] + if e == nil { + continue + } + if e.DailyWindowStart == nil || e.WeeklyWindowStart == nil || e.MonthlyWindowStart == nil { + continue + } + snaps = append(snaps, UserPlatformQuotaSnapshot{ + UserID: key.UserID, + Platform: key.Platform, + DailyUsageUSD: e.DailyUsageUSD, + WeeklyUsageUSD: e.WeeklyUsageUSD, + MonthlyUsageUSD: e.MonthlyUsageUSD, + DailyWindowStart: *e.DailyWindowStart, + WeeklyWindowStart: *e.WeeklyWindowStart, + MonthlyWindowStart: *e.MonthlyWindowStart, + }) + } + + // 4. 全部 MISS/异常跳过时 + if len(snaps) == 0 { + // 若 Pop 数量已不满一批,表示脏集将空,停止 + if len(keys) < s.batchSize { + return false + } + // 否则继续下一批(可能还有更多脏 key) + return true + } + + // 已知竞态(admin 写 × flusher 刷,仅 flusher_enabled=true 时存在): + // admin ResetExpiredWindow/UpsertForUser 是"先写 DB 再 DeleteCache"。若本批已 SPOP + BatchGet + // 读到旧 usage 快照(此刻 member 已离开脏集),而 admin 随后写 DB、本行 UPSERT 又在 admin 写之后落库, + // 则旧快照会覆盖 admin 刚写入的值;DeleteCache 后 Redis MISS,下次 preflight 从 DB 重载被覆盖的旧值。 + // 因 member 已被 SPOP,admin 侧 SREM/清脏标记无法拦截本批(故未做)。影响有限,暂列为已知取舍: + // - UpsertForUser 改 limit,而本 UPSERT 不写 limit 列 → limit 配置不受影响; + // - ResetExpiredWindow 改 usage,但 preflight windowExpired 会在窗口真正过期时自愈重置, + // 仅"强制重置未过期窗口"且与本批精确交错时短暂失效; + // - 低频 admin 操作 + 默认 flusher_enabled=false。彻底消除需 version OCC(DB 加 version 列条件 UPSERT), + // 成本高;启用 flusher 后如需强一致再评估。 + + // 5. 写入 DB + start := time.Now() + writeErr := s.quotaRepo.BatchSnapshotUsage(ctx, snaps, time.Now().UTC()) + s.updateLatencyMax(time.Since(start).Milliseconds()) + + if writeErr != nil { + if errors.Is(writeErr, ErrUserPlatformQuotaFKViolation) { + // 注意:PG FK violation 是整条 INSERT 回滚 → 整批(含同批正常用户)均未写入 DB, + // 且这些 key 已被 SPOP 出脏集、此处不 Readd。活跃 key 会在下次请求重新 SADD, + // flusher 读 Redis 当前累计绝对值刷库即自愈;低活跃 key 这轮 DB usage 偏低 + // (Redis 仍是 enforcement 权威,不受影响;DB 仅展示)。已删用户边角的接受取舍,不做逐行重试。 + // FK 违反:用户已被删除,直接丢弃不 Readd + s.metrics.FlushFKViolationTotal.Add(1) + s.metrics.FlushErrorTotal.Add(1) + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] FK violation (dropped %d snaps): %v", len(snaps), writeErr) + } else { + // 其他错误:回填脏集,保留下次重试 + s.metrics.FlushErrorTotal.Add(1) + s.readdOrCountLost(ctx, keys, "BatchSnapshotUsage") + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] BatchSnapshotUsage error: %v", writeErr) + } + return false + } + + // 6. 成功 + s.metrics.FlushSuccessTotal.Add(1) + s.metrics.FlushBatchSizeTotal.Add(int64(len(snaps))) + + // 若 Pop 数量不满一批,脏集已空,停止 + if len(keys) < s.batchSize { + return false + } + return true +} + +// flush 执行一次完整的 flush,循环消费至脏集空或达到 maxBatchesPerTick。 +func (s *UserPlatformQuotaUsageFlusher) flush() { + if s == nil { + return + } + parentCtx := context.Background() + for b := 0; b < flusherMaxBatchesPerTick; b++ { + if !s.flushOneBatch(parentCtx) { + return + } + } + // 连续消费满 flusherMaxBatchesPerTick 批仍未取空脏集:本 tick 主动让出,剩余积压留待下一 tick。 + // 记一条 log 便于 oncall 发现 distinct 活跃 key 远超 maxBatchesPerTick×batchSize(DB 镜像延迟上升); + // 可配合 Redis SCARD billing:upq:dirty 观察脏集存量。 + logger.LegacyPrintf("quota_flusher", + "[QuotaFlusher] 单 tick 达到 max batches 上限(%d × batchSize=%d),脏集仍非空,积压顺延至下一 tick", + flusherMaxBatchesPerTick, s.batchSize) +} + +// tick 是 TimingWheel 回调。若 flusher 已停止则直接返回。 +func (s *UserPlatformQuotaUsageFlusher) tick() { + if s == nil || s.stopped.Load() { + return + } + s.flush() +} + +// Start 注册定时 tick。flusher_enabled=false 时直接返回,不注册定时器。 +func (s *UserPlatformQuotaUsageFlusher) Start() { + if s == nil || !s.enabled { + return + } + s.timingWheel.ScheduleRecurring("deferred:platform_quota", s.interval, s.tick) +} + +// Stop 停止 flusher:标记 stopped → Cancel 定时器 → 执行最后一次 flush。 +func (s *UserPlatformQuotaUsageFlusher) Stop() { + if s == nil { + return + } + s.stopped.Store(true) + s.timingWheel.Cancel("deferred:platform_quota") + s.flush() +} diff --git a/backend/internal/service/user_platform_quota_flusher_test.go b/backend/internal/service/user_platform_quota_flusher_test.go new file mode 100644 index 00000000..4f734481 --- /dev/null +++ b/backend/internal/service/user_platform_quota_flusher_test.go @@ -0,0 +1,511 @@ +package service + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" +) + +// --------------------------------------------------------------------------- +// Mock: quotaDirtyCache +// --------------------------------------------------------------------------- + +type mockQuotaDirtyCache struct { + // popSequence: 第 0 次 Pop 返回 popSequence[0],之后返回 nil(空集) + popSequence [][]UserPlatformQuotaKey + popCallIdx int + + // getEntries: BatchGetUserPlatformQuotaCache 返回的 entries(与 keys 对齐) + getEntries []*UserPlatformQuotaCacheEntry + getErr error + + // readdCalled: 记录 Readd 收到的 keys(累积所有次调用) + readdCalled [][]UserPlatformQuotaKey + readdErr error +} + +func (m *mockQuotaDirtyCache) PopDirtyUserPlatformQuotaKeys(_ context.Context, _ int) ([]UserPlatformQuotaKey, error) { + if m.popCallIdx < len(m.popSequence) { + keys := m.popSequence[m.popCallIdx] + m.popCallIdx++ + return keys, nil + } + // 超出序列 → 空集(模拟脏集已清空) + return nil, nil +} + +func (m *mockQuotaDirtyCache) ReaddDirtyUserPlatformQuotaKeys(_ context.Context, keys []UserPlatformQuotaKey) error { + m.readdCalled = append(m.readdCalled, keys) + return m.readdErr +} + +func (m *mockQuotaDirtyCache) BatchGetUserPlatformQuotaCache(_ context.Context, _ []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + if m.getErr != nil { + return nil, m.getErr + } + return m.getEntries, nil +} + +// --------------------------------------------------------------------------- +// Mock: quotaSnapshotWriter +// --------------------------------------------------------------------------- + +type mockQuotaSnapshotWriter struct { + receivedSnaps []UserPlatformQuotaSnapshot + returnErr error +} + +func (m *mockQuotaSnapshotWriter) BatchSnapshotUsage(_ context.Context, snaps []UserPlatformQuotaSnapshot, _ time.Time) error { + m.receivedSnaps = append(m.receivedSnaps, snaps...) + return m.returnErr +} + +// --------------------------------------------------------------------------- +// Helper: 构造窗口起始时间(非 nil) +// --------------------------------------------------------------------------- + +func flusherPtrTime(t time.Time) *time.Time { return &t } + +func makeEntry(daily, weekly, monthly float64) *UserPlatformQuotaCacheEntry { + now := time.Now().UTC() + return &UserPlatformQuotaCacheEntry{ + DailyUsageUSD: daily, + WeeklyUsageUSD: weekly, + MonthlyUsageUSD: monthly, + DailyWindowStart: flusherPtrTime(now), + WeeklyWindowStart: flusherPtrTime(now), + MonthlyWindowStart: flusherPtrTime(now), + } +} + +// --------------------------------------------------------------------------- +// newTestFlusher: 直接构造 struct(跳过构造函数,B7 才注入) +// --------------------------------------------------------------------------- + +func newTestFlusher(cache quotaDirtyCache, writer quotaSnapshotWriter) *UserPlatformQuotaUsageFlusher { + return &UserPlatformQuotaUsageFlusher{ + cache: cache, + quotaRepo: writer, + timingWheel: nil, // 单测不启动 TimingWheel + interval: 5 * time.Second, + batchSize: 100, + flushTimeout: 5 * time.Second, + metrics: &FlusherMetrics{}, + } +} + +// --------------------------------------------------------------------------- +// 场景 1: PopSnapshotUpsert — 2 key + 2 个含 window 的 entry → writer 收 2 行 +// --------------------------------------------------------------------------- + +func TestFlusher_PopSnapshotUpsert(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 1, Platform: "anthropic"}, + {UserID: 2, Platform: "openai"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, // 第 1 次返回 keys,之后空 + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + makeEntry(4.0, 5.0, 6.0), + }, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + f.flush() + + if len(writer.receivedSnaps) != 2 { + t.Fatalf("expected 2 snaps, got %d", len(writer.receivedSnaps)) + } + if f.metrics.FlushBatchSizeTotal.Load() != 2 { + t.Errorf("FlushBatchSizeTotal = %d, want 2", f.metrics.FlushBatchSizeTotal.Load()) + } + if f.metrics.FlushSuccessTotal.Load() != 1 { + t.Errorf("FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load()) + } + if f.metrics.FlushErrorTotal.Load() != 0 { + t.Errorf("FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load()) + } +} + +// --------------------------------------------------------------------------- +// 场景 2: MissKeySkipped — 2 key,BatchGet 返回 [entry, nil] → 只刷 1 行,nil 跳过,不 Readd +// --------------------------------------------------------------------------- + +func TestFlusher_MissKeySkipped(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 1, Platform: "anthropic"}, + {UserID: 2, Platform: "openai"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + nil, // MISS + }, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + f.flush() + + if len(writer.receivedSnaps) != 1 { + t.Fatalf("expected 1 snap, got %d", len(writer.receivedSnaps)) + } + if writer.receivedSnaps[0].UserID != 1 { + t.Errorf("expected snap for UserID=1, got %d", writer.receivedSnaps[0].UserID) + } + if len(cache.readdCalled) != 0 { + t.Errorf("Readd should NOT be called on MISS, got %d calls", len(cache.readdCalled)) + } + if f.metrics.FlushSuccessTotal.Load() != 1 { + t.Errorf("FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load()) + } +} + +// --------------------------------------------------------------------------- +// 场景 3: UpsertFailReadds — writer 返普通 error → keys 被 Readd,FlushErrorTotal=1,DirtyReaddTotal=len +// --------------------------------------------------------------------------- + +func TestFlusher_UpsertFailReadds(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 1, Platform: "anthropic"}, + {UserID: 2, Platform: "openai"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + makeEntry(4.0, 5.0, 6.0), + }, + } + writeErr := errors.New("db connection timeout") + writer := &mockQuotaSnapshotWriter{returnErr: writeErr} + f := newTestFlusher(cache, writer) + + f.flush() + + if f.metrics.FlushErrorTotal.Load() != 1 { + t.Errorf("FlushErrorTotal = %d, want 1", f.metrics.FlushErrorTotal.Load()) + } + if len(cache.readdCalled) == 0 { + t.Fatal("Readd should be called after write error") + } + totalReadd := 0 + for _, rk := range cache.readdCalled { + totalReadd += len(rk) + } + if totalReadd != len(keys) { + t.Errorf("DirtyReaddTotal (from Readd calls) = %d, want %d", totalReadd, len(keys)) + } + if f.metrics.DirtyReaddTotal.Load() != int64(len(keys)) { + t.Errorf("DirtyReaddTotal metric = %d, want %d", f.metrics.DirtyReaddTotal.Load(), len(keys)) + } + if f.metrics.FlushSuccessTotal.Load() != 0 { + t.Errorf("FlushSuccessTotal = %d, want 0", f.metrics.FlushSuccessTotal.Load()) + } +} + +// --------------------------------------------------------------------------- +// 场景 4: FKViolationDropsNoReadd — writer 返 ErrUserPlatformQuotaFKViolation → 不 Readd,FlushFKViolationTotal=1 +// --------------------------------------------------------------------------- + +func TestFlusher_FKViolationDropsNoReadd(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 999, Platform: "anthropic"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + }, + } + writer := &mockQuotaSnapshotWriter{returnErr: ErrUserPlatformQuotaFKViolation} + f := newTestFlusher(cache, writer) + + f.flush() + + if f.metrics.FlushFKViolationTotal.Load() != 1 { + t.Errorf("FlushFKViolationTotal = %d, want 1", f.metrics.FlushFKViolationTotal.Load()) + } + if f.metrics.FlushErrorTotal.Load() != 1 { + t.Errorf("FlushErrorTotal = %d, want 1", f.metrics.FlushErrorTotal.Load()) + } + if len(cache.readdCalled) != 0 { + t.Errorf("Readd should NOT be called for FK violation (drop), got %d calls", len(cache.readdCalled)) + } + if f.metrics.DirtyReaddTotal.Load() != 0 { + t.Errorf("DirtyReaddTotal = %d, want 0 (FK violation drops)", f.metrics.DirtyReaddTotal.Load()) + } +} + +// --------------------------------------------------------------------------- +// 场景 5: NilSafe — var f *UserPlatformQuotaUsageFlusher; f.flush(); f.Stop() 不 panic +// --------------------------------------------------------------------------- + +func TestFlusher_NilSafe(t *testing.T) { + var f *UserPlatformQuotaUsageFlusher + // 下面两行不应 panic + f.flush() + f.Stop() +} + +// --------------------------------------------------------------------------- +// 场景 6: StopPreventsFlush — stopped=true 后 tick() 不调 flush(writer 没收到 snaps) +// --------------------------------------------------------------------------- + +func TestFlusher_StopPreventsFlush(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 1, Platform: "anthropic"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + }, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + // 标记为已停止 + f.stopped.Store(true) + + // tick 应该直接返回,不触发 flush + f.tick() + + if len(writer.receivedSnaps) != 0 { + t.Errorf("expected 0 snaps after stop, got %d", len(writer.receivedSnaps)) + } + if cache.popCallIdx != 0 { + t.Errorf("Pop should not be called after stop, popCallIdx = %d", cache.popCallIdx) + } +} + +// --------------------------------------------------------------------------- +// 场景 B13-1: ZeroPercentCompany — 0% 公司脏集恒空,flusher 空跑无 DB 写 +// +// 模拟几乎没有用户配置 quota limit 的公司:脏集始终为空(popSequence 为空切片), +// Pop 每次返回空集。flush() 应早退,不写 DB、不计成功、不 Readd。 +// --------------------------------------------------------------------------- + +func TestScenario_ZeroPercentCompany(t *testing.T) { + cache := &mockQuotaDirtyCache{ + // popSequence 为空 → Pop 超出序列 → 始终返回 nil(空集) + popSequence: [][]UserPlatformQuotaKey{}, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + f.flush() + + if len(writer.receivedSnaps) != 0 { + t.Errorf("0%% company: expected 0 snaps, got %d", len(writer.receivedSnaps)) + } + if f.metrics.FlushBatchSizeTotal.Load() != 0 { + t.Errorf("0%% company: FlushBatchSizeTotal = %d, want 0", f.metrics.FlushBatchSizeTotal.Load()) + } + if f.metrics.FlushSuccessTotal.Load() != 0 { + t.Errorf("0%% company: FlushSuccessTotal = %d, want 0 (empty-set early return)", f.metrics.FlushSuccessTotal.Load()) + } + if f.metrics.FlushErrorTotal.Load() != 0 { + t.Errorf("0%% company: FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load()) + } + if len(cache.readdCalled) != 0 { + t.Errorf("0%% company: Readd should never be called, got %d calls", len(cache.readdCalled)) + } +} + +// --------------------------------------------------------------------------- +// P1: IntervalFallback — flush_interval_ms ≤0 时回退 2s;正常值保留 +// --------------------------------------------------------------------------- + +func TestNewUserPlatformQuotaUsageFlusher_IntervalFallback(t *testing.T) { + cases := []struct { + name string + inMs int + wantDu time.Duration + }{ + {"零值回退 2s", 0, 2 * time.Second}, + {"负数回退 2s", -100, 2 * time.Second}, + {"正常 2000ms 保留", 2000, 2 * time.Second}, + {"正常 500ms 保留", 500, 500 * time.Millisecond}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cfg := &config.Config{} + cfg.Database.UserPlatformQuotaFlushIntervalMs = tc.inMs + f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil) + if f.interval != tc.wantDu { + t.Fatalf("interval = %v, want %v", f.interval, tc.wantDu) + } + }) + } +} + +// --------------------------------------------------------------------------- +// P1: EnabledField — flusher_enabled 配置正确写入 f.enabled +// --------------------------------------------------------------------------- + +func TestNewUserPlatformQuotaUsageFlusher_EnabledField(t *testing.T) { + for _, enabled := range []bool{true, false} { + cfg := &config.Config{} + cfg.Database.UserPlatformQuotaFlusherEnabled = enabled + cfg.Database.UserPlatformQuotaFlushIntervalMs = 500 + f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil) + if f.enabled != enabled { + t.Errorf("enabled = %v, want %v", f.enabled, enabled) + } + } +} + +// --------------------------------------------------------------------------- +// P2: ReaddFailCounts — BatchGet 失败 + Readd 失败 → DirtyLostTotal 增、DirtyReaddTotal 不变 +// BatchGet 失败 + Readd 成功 → DirtyReaddTotal 增、DirtyLostTotal 不变 +// --------------------------------------------------------------------------- + +func TestFlusher_ReaddFailCounts(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 10, Platform: "anthropic"}, + {UserID: 11, Platform: "openai"}, + } + + t.Run("Readd 失败计 DirtyLostTotal", func(t *testing.T) { + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getErr: errors.New("redis timeout"), // 触发 BatchGet 失败路径 + readdErr: errors.New("redis connection refused"), // Readd 也失败 + } + f := newTestFlusher(cache, &mockQuotaSnapshotWriter{}) + + f.flush() + + if f.metrics.DirtyLostTotal.Load() != int64(len(keys)) { + t.Errorf("DirtyLostTotal = %d, want %d", f.metrics.DirtyLostTotal.Load(), len(keys)) + } + if f.metrics.DirtyReaddTotal.Load() != 0 { + t.Errorf("DirtyReaddTotal = %d, want 0 (Readd 失败不应计入)", f.metrics.DirtyReaddTotal.Load()) + } + }) + + t.Run("Readd 成功计 DirtyReaddTotal", func(t *testing.T) { + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getErr: errors.New("redis timeout"), // 触发 BatchGet 失败路径 + readdErr: nil, // Readd 成功 + } + f := newTestFlusher(cache, &mockQuotaSnapshotWriter{}) + + f.flush() + + if f.metrics.DirtyReaddTotal.Load() != int64(len(keys)) { + t.Errorf("DirtyReaddTotal = %d, want %d", f.metrics.DirtyReaddTotal.Load(), len(keys)) + } + if f.metrics.DirtyLostTotal.Load() != 0 { + t.Errorf("DirtyLostTotal = %d, want 0 (Readd 成功不应计 lost)", f.metrics.DirtyLostTotal.Load()) + } + }) +} + +// --------------------------------------------------------------------------- +// ClampsBatchSize — NewUserPlatformQuotaUsageFlusher 构造时按 +// [defaultFlushBatchSize, maxFlushBatchSize] 区间 clamp batchSize +// --------------------------------------------------------------------------- + +func TestNewUserPlatformQuotaUsageFlusher_ClampsBatchSize(t *testing.T) { + cases := []struct { + name string + in int + want int + }{ + {"超上限被 clamp", 7000, maxFlushBatchSize}, + {"恰好上限保留", maxFlushBatchSize, maxFlushBatchSize}, + {"零回退默认", 0, defaultFlushBatchSize}, + {"负数回退默认", -5, defaultFlushBatchSize}, + {"正常值保留", 500, 500}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cfg := &config.Config{} + cfg.Database.UserPlatformQuotaFlushBatchSize = tc.in + f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil) + if f.batchSize != tc.want { + t.Fatalf("batchSize = %d, want %d", f.batchSize, tc.want) + } + }) + } +} + +// --------------------------------------------------------------------------- +// 场景 B13-2: NinetyPercentCompany — 90% 公司大量用户配 limit,一批 5 key 批量刷库 +// +// 模拟大量用户配置了 quota limit 的公司:脏集第一次 Pop 返回 5 个不同用户的 key, +// 之后返回空集(避免 flush 循环)。flush() 应构造 5 条 snapshot 写入 DB, +// 断言绝对值语义(snap 的 DailyUsageUSD 等于 entry 的值)、metrics 正确、不 Readd。 +// --------------------------------------------------------------------------- + +func TestScenario_NinetyPercentCompany(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 101, Platform: "anthropic"}, + {UserID: 102, Platform: "anthropic"}, + {UserID: 103, Platform: "openai"}, + {UserID: 104, Platform: "openai"}, + {UserID: 105, Platform: "anthropic"}, + } + entries := []*UserPlatformQuotaCacheEntry{ + makeEntry(1.1, 2.2, 3.3), + makeEntry(4.4, 5.5, 6.6), + makeEntry(7.7, 8.8, 9.9), + makeEntry(0.5, 1.0, 1.5), + makeEntry(10.0, 20.0, 30.0), + } + cache := &mockQuotaDirtyCache{ + // 第 1 次 Pop 返回 5 keys,之后返回空集(防止 flush 无限循环) + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: entries, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + f.flush() + + // 应收到 5 条 snapshot + if len(writer.receivedSnaps) != 5 { + t.Fatalf("90%% company: expected 5 snaps, got %d", len(writer.receivedSnaps)) + } + + // 验证绝对值语义:第 1 条 snap 的各窗口 usage 应等于 entries[0] 的值 + snap0 := writer.receivedSnaps[0] + entry0 := entries[0] + if snap0.DailyUsageUSD != entry0.DailyUsageUSD { + t.Errorf("snap[0].DailyUsageUSD = %v, want %v", snap0.DailyUsageUSD, entry0.DailyUsageUSD) + } + if snap0.WeeklyUsageUSD != entry0.WeeklyUsageUSD { + t.Errorf("snap[0].WeeklyUsageUSD = %v, want %v", snap0.WeeklyUsageUSD, entry0.WeeklyUsageUSD) + } + if snap0.MonthlyUsageUSD != entry0.MonthlyUsageUSD { + t.Errorf("snap[0].MonthlyUsageUSD = %v, want %v", snap0.MonthlyUsageUSD, entry0.MonthlyUsageUSD) + } + + // FlushBatchSizeTotal 应为 5(本批 keys 数量) + if f.metrics.FlushBatchSizeTotal.Load() != 5 { + t.Errorf("90%% company: FlushBatchSizeTotal = %d, want 5", f.metrics.FlushBatchSizeTotal.Load()) + } + // FlushSuccessTotal 应为 1(1 个批次写成功) + if f.metrics.FlushSuccessTotal.Load() != 1 { + t.Errorf("90%% company: FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load()) + } + // 无错误、无 Readd + if f.metrics.FlushErrorTotal.Load() != 0 { + t.Errorf("90%% company: FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load()) + } + if f.metrics.DirtyReaddTotal.Load() != 0 { + t.Errorf("90%% company: DirtyReaddTotal = %d, want 0", f.metrics.DirtyReaddTotal.Load()) + } + if len(cache.readdCalled) != 0 { + t.Errorf("90%% company: Readd should not be called, got %d calls", len(cache.readdCalled)) + } +} diff --git a/backend/internal/service/user_platform_quota_port.go b/backend/internal/service/user_platform_quota_port.go index cb09542a..0f88eda4 100644 --- a/backend/internal/service/user_platform_quota_port.go +++ b/backend/internal/service/user_platform_quota_port.go @@ -11,6 +11,23 @@ import ( // handler 只需引用 service 包,无需直接依赖 repository 包。 var ErrUserPlatformQuotaNotFound = errors.New("user platform quota not found") +// ErrUserPlatformQuotaFKViolation service 层 sentinel:批量 snapshot UPSERT 时存在 +// user_id 不在 users 表的记录(外键违反)。adapter 负责将 repository 层同名 sentinel 包装为此错误。 +var ErrUserPlatformQuotaFKViolation = errors.New("user platform quota snapshot FK violation") + +// UserPlatformQuotaSnapshot 是 service 层 flusher 向 DB 写入快照时使用的传输结构。 +// 字段语义与 repository.UserPlatformQuotaSnapshot 完全对应,由 adapter 负责转换。 +type UserPlatformQuotaSnapshot struct { + UserID int64 + Platform string + DailyUsageUSD float64 + WeeklyUsageUSD float64 + MonthlyUsageUSD float64 + DailyWindowStart time.Time + WeeklyWindowStart time.Time + MonthlyWindowStart time.Time +} + // UserPlatformQuotaRecord service 层传输结构体(与 repository 层解耦)。 type UserPlatformQuotaRecord struct { UserID int64 @@ -47,4 +64,6 @@ type UserPlatformQuotaRepository interface { // ResetExpiredWindow 重置指定窗口("daily"|"weekly"|"monthly")的用量与起始时间。 // 未命中活跃记录时返回(service-side wrapper of repository.ErrUserPlatformQuotaNotFound)。 ResetExpiredWindow(ctx context.Context, userID int64, platform string, window string, newStart time.Time) error + // BatchSnapshotUsage 绝对值覆盖写入整批 usage 快照。FK 违反返回 ErrUserPlatformQuotaFKViolation。 + BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error } diff --git a/backend/internal/service/user_service_test.go b/backend/internal/service/user_service_test.go index 19aec5d3..1a18e70a 100644 --- a/backend/internal/service/user_service_test.go +++ b/backend/internal/service/user_service_test.go @@ -327,10 +327,22 @@ func (m *mockBillingCache) DeleteUserPlatformQuotaCache(context.Context, int64, return nil } -func (m *mockBillingCache) IncrUserPlatformQuotaUsageCache(context.Context, int64, string, float64, time.Duration) error { +func (m *mockBillingCache) IncrUserPlatformQuotaUsageCache(context.Context, int64, string, float64, time.Duration, bool) error { return nil } +func (m *mockBillingCache) PopDirtyUserPlatformQuotaKeys(context.Context, int) ([]UserPlatformQuotaKey, error) { + return nil, nil +} + +func (m *mockBillingCache) ReaddDirtyUserPlatformQuotaKeys(context.Context, []UserPlatformQuotaKey) error { + return nil +} + +func (m *mockBillingCache) BatchGetUserPlatformQuotaCache(context.Context, []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + return nil, nil +} + // --- 测试 --- func TestUpdateBalance_Success(t *testing.T) { diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index d3e4ce51..19bd841d 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -571,8 +571,16 @@ var ProviderSet = wire.NewSet( ProvideChannelMonitorService, ProvideChannelMonitorRunner, NewChannelMonitorRequestTemplateService, + ProvideUserPlatformQuotaUsageFlusher, ) +// ProvideUserPlatformQuotaUsageFlusher 创建并启动 UserPlatformQuotaUsageFlusher。 +func ProvideUserPlatformQuotaUsageFlusher(cfg *config.Config, cache BillingCache, quotaRepo UserPlatformQuotaRepository, tw *TimingWheelService) *UserPlatformQuotaUsageFlusher { + svc := NewUserPlatformQuotaUsageFlusher(cfg, cache, quotaRepo, tw) + svc.Start() + return svc +} + // ProvidePaymentConfigService wraps NewPaymentConfigService to accept the named // payment.EncryptionKey type instead of raw []byte, avoiding Wire ambiguity. func ProvidePaymentConfigService(entClient *dbent.Client, settingRepo SettingRepository, key payment.EncryptionKey) *PaymentConfigService {