Files
app-api/internal/cache/share_code_cache.go
T
eafonyang b2aed2a6b5 feat(share): 新增 shareList 接口 + 分享码模块全链路优化
- 新增 GET /audio/shareList?mac=xx 查询未过期分享码(纯 Redis,ZSET 索引)
- share:mac:{mac} 从 SET 改为 ZSET(score=expire_at),支持 ZREMRANGEBYSCORE 精确过期清理
- 空 ZSET 自动 DEL,避免 key 累积
- share:import:pending 增加 12h 兜底 TTL,防止 DB 不可用时内存泄漏
- 导入日志 field 改为 mac:code(去掉 nanotime),同 MAC+code 多次导入幂等去重
- MarkPersisted 保存/恢复 PTTL,防御性编程
- shareCreate 改为 POST + JSON body
- 全量 handler 补充 Swagger 注释,集成 swag 文档生成
- Makefile 使用 $(go env GOPATH)/bin/swag 解决 PATH 问题
2026-06-12 17:07:59 +08:00

279 lines
7.6 KiB
Go

package cache
import (
"context"
"crypto/rand"
"encoding/json"
"errors"
"fmt"
"math/big"
"strings"
"time"
"github.com/redis/go-redis/v9"
)
const (
shareCodeKeyPrefix = "share:"
sharePendingSet = "share:pending"
shareImportPendingHash = "share:import:pending"
shareFlushLockPref = "share:flush:lock:"
shareImportFlushLockPref = "share:import:flush:lock:"
shareMacIndexPrefix = "share:mac:"
shareCodeLength = 5
shareCodeTTL = 30 * time.Minute
shareImportPendingTTL = 12 * time.Hour
shareCodeCharset = "23456789ABCDEFGHJKLMNPQRSTUVWXYZ"
shareCodeMaxRetries = 20
fieldMacAddr = "mac_addr"
fieldIPAddr = "ip_addr"
fieldEqData = "eq_data"
fieldExpireAt = "expire_at"
fieldPersisted = "persisted"
)
var shareCreateScript = redis.NewScript(`
if redis.call('EXISTS', KEYS[1]) == 1 then
return 0
end
redis.call('HSET', KEYS[1],
'mac_addr', ARGV[1],
'ip_addr', ARGV[2],
'eq_data', ARGV[3],
'expire_at', ARGV[4],
'persisted', '0'
)
redis.call('EXPIRE', KEYS[1], tonumber(ARGV[5]))
redis.call('SADD', KEYS[2], ARGV[6])
-- ZSET: score = expire_at unix timestamp, member = share code
redis.call('ZADD', KEYS[3], tonumber(ARGV[7]), ARGV[6])
return 1
`)
type ShareImportPendingLog struct {
MacAddr string `json:"mac_addr"`
ShareCode string `json:"share_code"`
IpAddr string `json:"ip_addr"`
EqData string `json:"eq_data"`
ExpireAt string `json:"expire_at"`
}
type ShareCodeData struct {
ShareCode string
MacAddr string
IpAddr string
EqData string
ExpireAt time.Time
Persisted bool
}
type ShareCodeCache struct {
rdb *redis.Client
}
func NewShareCodeCache(rdb *redis.Client) *ShareCodeCache {
return &ShareCodeCache{rdb: rdb}
}
func ShareCodeTTL() time.Duration {
return shareCodeTTL
}
func (c *ShareCodeCache) Create(ctx context.Context, macAddr, ipAddr string, eqData []byte) (*ShareCodeData, error) {
expireAt := time.Now().Add(shareCodeTTL)
eqJSON := string(eqData)
for i := 0; i < shareCodeMaxRetries; i++ {
code, err := randomShareCode()
if err != nil {
return nil, err
}
key := shareCodeKey(code)
macIdxKey := shareMacIndexKey(macAddr)
ok, err := shareCreateScript.Run(ctx, c.rdb, []string{key, sharePendingSet, macIdxKey},
macAddr, ipAddr, eqJSON, expireAt.Format(time.RFC3339), int(shareCodeTTL.Seconds()), code, expireAt.Unix(),
).Int()
if err != nil {
return nil, fmt.Errorf("create share code in redis: %w", err)
}
if ok == 1 {
return &ShareCodeData{
ShareCode: code,
MacAddr: macAddr,
IpAddr: ipAddr,
EqData: eqJSON,
ExpireAt: expireAt,
Persisted: false,
}, nil
}
}
return nil, errors.New("failed to allocate unique share code")
}
func (c *ShareCodeCache) Get(ctx context.Context, shareCode string) (*ShareCodeData, error) {
key := shareCodeKey(shareCode)
values, err := c.rdb.HGetAll(ctx, key).Result()
if err != nil {
return nil, fmt.Errorf("hgetall share code: %w", err)
}
if len(values) == 0 {
return nil, nil
}
expireAt, err := time.Parse(time.RFC3339, values[fieldExpireAt])
if err != nil {
return nil, fmt.Errorf("parse share code expire_at: %w", err)
}
return &ShareCodeData{
ShareCode: shareCode,
MacAddr: values[fieldMacAddr],
IpAddr: values[fieldIPAddr],
EqData: values[fieldEqData],
ExpireAt: expireAt,
Persisted: values[fieldPersisted] == "1",
}, nil
}
func (c *ShareCodeCache) ListPending(ctx context.Context) ([]string, error) {
codes, err := c.rdb.SMembers(ctx, sharePendingSet).Result()
if err != nil {
return nil, fmt.Errorf("smembers share pending: %w", err)
}
return codes, nil
}
func (c *ShareCodeCache) AcquireFlushLock(ctx context.Context, shareCode string) (bool, error) {
lockKey := shareFlushLockPref + shareCode
return c.rdb.SetNX(ctx, lockKey, "1", 30*time.Second).Result()
}
func (c *ShareCodeCache) ReleaseFlushLock(ctx context.Context, shareCode string) error {
return c.rdb.Del(ctx, shareFlushLockPref+shareCode).Err()
}
func (c *ShareCodeCache) MarkPersisted(ctx context.Context, shareCode string) error {
key := shareCodeKey(shareCode)
// preserve remaining TTL before HSET (HSET clears TTL in some Redis versions)
remainTTL, err := c.rdb.PTTL(ctx, key).Result()
if err != nil {
return fmt.Errorf("get share code ttl: %w", err)
}
if err := c.rdb.HSet(ctx, key, fieldPersisted, "1").Err(); err != nil {
return fmt.Errorf("mark share code persisted: %w", err)
}
// restore TTL
if remainTTL > 0 {
_ = c.rdb.PExpire(ctx, key, remainTTL).Err()
}
return c.rdb.SRem(ctx, sharePendingSet, shareCode).Err()
}
func (c *ShareCodeCache) EnqueueImportLog(ctx context.Context, macAddr, shareCode, ipAddr, eqData string, expireAt time.Time) error {
field := fmt.Sprintf("%s:%s", macAddr, shareCode) // 同 MAC+code 幂等,避免重复导入产生多条记录
payload, err := json.Marshal(ShareImportPendingLog{
MacAddr: macAddr,
ShareCode: shareCode,
IpAddr: ipAddr,
EqData: eqData,
ExpireAt: expireAt.Format(time.RFC3339),
})
if err != nil {
return fmt.Errorf("marshal import pending log: %w", err)
}
if err := c.rdb.HSet(ctx, shareImportPendingHash, field, payload).Err(); err != nil {
return fmt.Errorf("enqueue import log: %w", err)
}
// Fallback TTL: prevent unbounded growth if persist task is disabled or DB is down
c.rdb.Expire(ctx, shareImportPendingHash, shareImportPendingTTL)
return nil
}
func (c *ShareCodeCache) ListPendingImports(ctx context.Context) (map[string]string, error) {
logs, err := c.rdb.HGetAll(ctx, shareImportPendingHash).Result()
if err != nil {
return nil, fmt.Errorf("hgetall import pending: %w", err)
}
return logs, nil
}
func (c *ShareCodeCache) RemovePendingImport(ctx context.Context, field string) error {
return c.rdb.HDel(ctx, shareImportPendingHash, field).Err()
}
func (c *ShareCodeCache) AcquireImportFlushLock(ctx context.Context, field string) (bool, error) {
return c.rdb.SetNX(ctx, shareImportFlushLockPref+field, "1", 30*time.Second).Result()
}
func (c *ShareCodeCache) ReleaseImportFlushLock(ctx context.Context, field string) error {
return c.rdb.Del(ctx, shareImportFlushLockPref+field).Err()
}
func (c *ShareCodeCache) ListByMac(ctx context.Context, macAddr string) ([]*ShareCodeData, error) {
key := shareMacIndexKey(macAddr)
now := time.Now()
// Remove expired entries from ZSET (score = expire_at unix timestamp)
_, err := c.rdb.ZRemRangeByScore(ctx, key, "-inf", fmt.Sprintf("%d", now.Unix())).Result()
if err != nil && strings.Contains(err.Error(), "WRONGTYPE") {
// Old SET-type key from previous version, delete it
_ = c.rdb.Del(ctx, key).Err()
return nil, nil
}
// Get remaining (unexpired) codes
codes, err := c.rdb.ZRange(ctx, key, 0, -1).Result()
if err != nil {
return nil, fmt.Errorf("zrange mac index: %w", err)
}
if len(codes) == 0 {
// ZSET is empty, remove the key to avoid accumulating empty keys
_ = c.rdb.Del(ctx, key).Err()
return nil, nil
}
var result []*ShareCodeData
for _, code := range codes {
data, err := c.Get(ctx, code)
if err != nil {
return nil, err
}
if data == nil {
// hash already expired/removed, clean from ZSET
_ = c.rdb.ZRem(ctx, key, code).Err()
continue
}
result = append(result, data)
}
return result, nil
}
func shareMacIndexKey(macAddr string) string {
return shareMacIndexPrefix + macAddr
}
func shareCodeKey(code string) string {
return shareCodeKeyPrefix + code
}
func randomShareCode() (string, error) {
b := make([]byte, shareCodeLength)
max := big.NewInt(int64(len(shareCodeCharset)))
for i := range b {
n, err := rand.Int(rand.Reader, max)
if err != nil {
return "", err
}
b[i] = shareCodeCharset[n.Int64()]
}
return string(b), nil
}