package cache import ( "context" "crypto/rand" "encoding/json" "errors" "fmt" "math/big" "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:" shareCodeLength = 5 shareCodeTTL = 30 * time.Minute 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]) 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) ok, err := shareCreateScript.Run(ctx, c.rdb, []string{key, sharePendingSet}, macAddr, ipAddr, eqJSON, expireAt.Format(time.RFC3339), int(shareCodeTTL.Seconds()), code, ).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) if err := c.rdb.HSet(ctx, key, fieldPersisted, "1").Err(); err != nil { return fmt.Errorf("mark share code persisted: %w", 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:%d", macAddr, shareCode, time.Now().UnixNano()) 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) } 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 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 }