77 "log"
88 "math/rand/v2"
99 "strconv"
10+ "strings"
1011 "time"
1112
1213 "github.com/Wei-Shaw/sub2api/internal/service"
@@ -338,38 +339,26 @@ func userPlatformQuotaCacheKey(userID int64, platform string) string {
338339 return fmt .Sprintf ("billing:user_platform_quota:%d:%s" , userID , platform )
339340}
340341
341- func (c * billingCache ) GetUserPlatformQuotaCache (ctx context.Context , userID int64 , platform string ) (* service.UserPlatformQuotaCacheEntry , bool , error ) {
342- key := userPlatformQuotaCacheKey (userID , platform )
343- fields := []string {
344- "daily_usage" , "weekly_usage" , "monthly_usage" , "version" , "schema_version" ,
345- "daily_limit" , "weekly_limit" , "monthly_limit" ,
346- "daily_window_start" , "weekly_window_start" , "monthly_window_start" ,
347- }
348- vals , err := c .rdb .HMGet (ctx , key , fields ... ).Result ()
349- if err != nil {
350- return nil , false , err
351- }
352- // 前4个全为nil → key 不存在
353- if vals [0 ] == nil && vals [1 ] == nil && vals [2 ] == nil && vals [3 ] == nil {
354- return nil , false , nil
342+ // parseUserPlatformQuotaHash 将 Redis HGETALL 返回的 map[string]string 反序列化为
343+ // *service.UserPlatformQuotaCacheEntry。空 map(key 不存在)返回 nil。
344+ // GetUserPlatformQuotaCache 和 BatchGetUserPlatformQuotaCache 共用此函数,确保解析逻辑一致。
345+ func parseUserPlatformQuotaHash (m map [string ]string ) * service.UserPlatformQuotaCacheEntry {
346+ if len (m ) == 0 {
347+ return nil
355348 }
356- parseFloat := func (v any ) float64 {
357- if v == nil {
349+ parseFloat := func (s string ) float64 {
350+ if s == "" {
358351 return 0
359352 }
360- s , ok := v .(string )
361- if ! ok {
353+ f , err := strconv .ParseFloat (s , 64 )
354+ if err != nil {
355+ log .Printf ("billing_cache: corrupt quota usage field %q (using 0): %v" , s , err )
362356 return 0
363357 }
364- f , _ := strconv .ParseFloat (s , 64 )
365358 return f
366359 }
367- parseFloatPtr := func (v any ) * float64 {
368- if v == nil {
369- return nil
370- }
371- s , ok := v .(string )
372- if ! ok || s == "" {
360+ parseFloatPtr := func (s string ) * float64 {
361+ if s == "" {
373362 return nil
374363 }
375364 f , err := strconv .ParseFloat (s , 64 )
@@ -378,12 +367,8 @@ func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int
378367 }
379368 return & f
380369 }
381- parseTimePtr := func (v any ) * time.Time {
382- if v == nil {
383- return nil
384- }
385- s , ok := v .(string )
386- if ! ok || s == "" {
370+ parseTimePtr := func (s string ) * time.Time {
371+ if s == "" {
387372 return nil
388373 }
389374 n , err := strconv .ParseInt (s , 10 , 64 )
@@ -393,30 +378,37 @@ func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int
393378 t := time .Unix (n , 0 ).UTC ()
394379 return & t
395380 }
396- parseInt64 := func (v any ) int64 {
397- if v == nil {
398- return 0
399- }
400- s , ok := v .(string )
401- if ! ok {
402- return 0
403- }
381+ parseInt64 := func (s string ) int64 {
404382 n , _ := strconv .ParseInt (s , 10 , 64 )
405383 return n
406384 }
407385 return & service.UserPlatformQuotaCacheEntry {
408- DailyUsageUSD : parseFloat (vals [0 ]),
409- WeeklyUsageUSD : parseFloat (vals [1 ]),
410- MonthlyUsageUSD : parseFloat (vals [2 ]),
411- Version : parseInt64 (vals [3 ]),
412- SchemaVersion : parseInt64 (vals [4 ]),
413- DailyLimitUSD : parseFloatPtr (vals [5 ]),
414- WeeklyLimitUSD : parseFloatPtr (vals [6 ]),
415- MonthlyLimitUSD : parseFloatPtr (vals [7 ]),
416- DailyWindowStart : parseTimePtr (vals [8 ]),
417- WeeklyWindowStart : parseTimePtr (vals [9 ]),
418- MonthlyWindowStart : parseTimePtr (vals [10 ]),
419- }, true , nil
386+ DailyUsageUSD : parseFloat (m ["daily_usage" ]),
387+ WeeklyUsageUSD : parseFloat (m ["weekly_usage" ]),
388+ MonthlyUsageUSD : parseFloat (m ["monthly_usage" ]),
389+ Version : parseInt64 (m ["version" ]),
390+ SchemaVersion : parseInt64 (m ["schema_version" ]),
391+ DailyLimitUSD : parseFloatPtr (m ["daily_limit" ]),
392+ WeeklyLimitUSD : parseFloatPtr (m ["weekly_limit" ]),
393+ MonthlyLimitUSD : parseFloatPtr (m ["monthly_limit" ]),
394+ DailyWindowStart : parseTimePtr (m ["daily_window_start" ]),
395+ WeeklyWindowStart : parseTimePtr (m ["weekly_window_start" ]),
396+ MonthlyWindowStart : parseTimePtr (m ["monthly_window_start" ]),
397+ }
398+ }
399+
400+ func (c * billingCache ) GetUserPlatformQuotaCache (ctx context.Context , userID int64 , platform string ) (* service.UserPlatformQuotaCacheEntry , bool , error ) {
401+ key := userPlatformQuotaCacheKey (userID , platform )
402+ m , err := c .rdb .HGetAll (ctx , key ).Result ()
403+ if err != nil {
404+ return nil , false , err
405+ }
406+ entry := parseUserPlatformQuotaHash (m )
407+ if entry == nil {
408+ // 空 map → key 不存在 → MISS
409+ return nil , false , nil
410+ }
411+ return entry , true , nil
420412}
421413
422414func (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
468460// SetCache 重建为新版 entry —— 若此处仍累加,上层覆盖时会丢失这部分增量,导致 Redis usage 比真实偏小。
469461// key 不存在同样跳过(由下次 SetCache 重建)。
470462// KEYS[1] = hash key
463+ // KEYS[2] = 脏集 key(dirty set)
471464// ARGV[1] = cost (string float)
472465// ARGV[2] = ttl seconds
473466// ARGV[3] = expected schema_version (Go 侧 UserPlatformQuotaCacheSchemaV1)
467+ // ARGV[4] = dirty set member(空串则不 SADD)
468+ // ARGV[5] = 脏集兜底 TTL 秒
474469const updateUserPlatformQuotaUsageScript = `
475470if redis.call("EXISTS", KEYS[1]) == 0 then
476471 return 0
@@ -484,18 +479,125 @@ redis.call("HINCRBYFLOAT", KEYS[1], "weekly_usage", ARGV[1])
484479redis.call("HINCRBYFLOAT", KEYS[1], "monthly_usage", ARGV[1])
485480redis.call("HINCRBY", KEYS[1], "version", 1)
486481redis.call("EXPIRE", KEYS[1], ARGV[2])
482+ if ARGV[4] ~= "" then
483+ redis.call("SADD", KEYS[2], ARGV[4])
484+ redis.call("EXPIRE", KEYS[2], ARGV[5])
485+ end
487486return 1
488487`
489488
490- func (c * billingCache ) IncrUserPlatformQuotaUsageCache (ctx context.Context , userID int64 , platform string , cost float64 , ttl time.Duration ) error {
491- key := userPlatformQuotaCacheKey (userID , platform )
492- _ , err := c .rdb .Eval (ctx , updateUserPlatformQuotaUsageScript , []string {key },
489+ // userPlatformQuotaDirtySetKey 返回脏集(dirty set)的 Redis key。
490+ // 使用与 userPlatformQuotaCacheKey 相同的前缀 "billing:"。
491+ func userPlatformQuotaDirtySetKey () string { return "billing:" + "upq:dirty" }
492+
493+ // userPlatformQuotaDirtyTTLSeconds 脏集兜底 TTL(秒):初始 SADD(Lua)与 Readd 共用,
494+ // 确保 flusher 长期停摆时脏集最终过期;正常运行因持续 SADD 不断续期。
495+ const userPlatformQuotaDirtyTTLSeconds = 86400
496+
497+ // userPlatformQuotaDirtyMember 构造脏集成员字符串 "userID:platform"。
498+ func userPlatformQuotaDirtyMember (userID int64 , platform string ) string {
499+ return strconv .FormatInt (userID , 10 ) + ":" + platform
500+ }
501+
502+ func (c * billingCache ) IncrUserPlatformQuotaUsageCache (ctx context.Context , userID int64 , platform string , cost float64 , ttl time.Duration , markDirty bool ) error {
503+ member := ""
504+ if markDirty {
505+ member = userPlatformQuotaDirtyMember (userID , platform )
506+ }
507+ _ , err := c .rdb .Eval (ctx , updateUserPlatformQuotaUsageScript ,
508+ []string {userPlatformQuotaCacheKey (userID , platform ), userPlatformQuotaDirtySetKey ()},
493509 strconv .FormatFloat (cost , 'f' , - 1 , 64 ),
494510 int (ttl .Seconds ()),
495511 service .UserPlatformQuotaCacheSchemaV1 ,
512+ member ,
513+ userPlatformQuotaDirtyTTLSeconds ,
496514 ).Result ()
497515 if err != nil && ! errors .Is (err , redis .Nil ) {
498516 return err
499517 }
500518 return nil
501519}
520+
521+ // parseUserPlatformQuotaDirtyMember 将脏集成员字符串 "userID:platform" 解析为
522+ // service.UserPlatformQuotaKey。解析失败返回 ok=false。
523+ func parseUserPlatformQuotaDirtyMember (m string ) (service.UserPlatformQuotaKey , bool ) {
524+ parts := strings .SplitN (m , ":" , 2 )
525+ if len (parts ) != 2 {
526+ return service.UserPlatformQuotaKey {}, false
527+ }
528+ uid , err := strconv .ParseInt (parts [0 ], 10 , 64 )
529+ if err != nil {
530+ return service.UserPlatformQuotaKey {}, false
531+ }
532+ return service.UserPlatformQuotaKey {UserID : uid , Platform : parts [1 ]}, true
533+ }
534+
535+ // PopDirtyUserPlatformQuotaKeys 从脏集随机弹出最多 n 个 key。
536+ // 脏集为空时返回 (nil, nil)。
537+ func (c * billingCache ) PopDirtyUserPlatformQuotaKeys (ctx context.Context , n int ) ([]service.UserPlatformQuotaKey , error ) {
538+ members , err := c .rdb .SPopN (ctx , userPlatformQuotaDirtySetKey (), int64 (n )).Result ()
539+ if err != nil {
540+ if errors .Is (err , redis .Nil ) {
541+ return nil , nil
542+ }
543+ return nil , err
544+ }
545+ keys := make ([]service.UserPlatformQuotaKey , 0 , len (members ))
546+ for _ , m := range members {
547+ k , ok := parseUserPlatformQuotaDirtyMember (m )
548+ if ! ok {
549+ log .Printf ("billing_cache: skipping invalid dirty member %q" , m )
550+ continue
551+ }
552+ keys = append (keys , k )
553+ }
554+ return keys , nil
555+ }
556+
557+ // ReaddDirtyUserPlatformQuotaKeys 将 keys 重新加入脏集(flush 失败时回填)。
558+ // 通过 pipeline 同时执行 SAdd + Expire,确保 Readd 后脏集具有兜底 TTL。
559+ // 空切片时直接返回 nil。
560+ func (c * billingCache ) ReaddDirtyUserPlatformQuotaKeys (ctx context.Context , keys []service.UserPlatformQuotaKey ) error {
561+ if len (keys ) == 0 {
562+ return nil
563+ }
564+ dirtyKey := userPlatformQuotaDirtySetKey ()
565+ members := make ([]any , len (keys ))
566+ for i , k := range keys {
567+ members [i ] = userPlatformQuotaDirtyMember (k .UserID , k .Platform )
568+ }
569+ pipe := c .rdb .Pipeline ()
570+ pipe .SAdd (ctx , dirtyKey , members ... )
571+ pipe .Expire (ctx , dirtyKey , userPlatformQuotaDirtyTTLSeconds * time .Second )
572+ _ , err := pipe .Exec (ctx )
573+ return err
574+ }
575+
576+ // BatchGetUserPlatformQuotaCache 通过 Pipeline 批量 HGETALL 获取多个 user×platform 的
577+ // quota cache。返回切片与 keys 顺序、长度对齐;MISS 或解析失败位置返回 nil。
578+ func (c * billingCache ) BatchGetUserPlatformQuotaCache (ctx context.Context , keys []service.UserPlatformQuotaKey ) ([]* service.UserPlatformQuotaCacheEntry , error ) {
579+ if len (keys ) == 0 {
580+ return nil , nil
581+ }
582+ pipe := c .rdb .Pipeline ()
583+ cmds := make ([]* redis.MapStringStringCmd , len (keys ))
584+ for i , k := range keys {
585+ cmds [i ] = pipe .HGetAll (ctx , userPlatformQuotaCacheKey (k .UserID , k .Platform ))
586+ }
587+ if _ , err := pipe .Exec (ctx ); err != nil && ! errors .Is (err , redis .Nil ) {
588+ return nil , err
589+ }
590+ results := make ([]* service.UserPlatformQuotaCacheEntry , len (keys ))
591+ for i , cmd := range cmds {
592+ m , err := cmd .Result ()
593+ if err != nil {
594+ if ! errors .Is (err , redis .Nil ) {
595+ log .Printf ("billing_cache: BatchGet HGETALL cmd[%d] failed: %v (skip, self-heal)" , i , err )
596+ }
597+ // 单个命令失败 → 对应位置 nil,继续
598+ continue
599+ }
600+ results [i ] = parseUserPlatformQuotaHash (m )
601+ }
602+ return results , nil
603+ }
0 commit comments