新增分享码功能

This commit is contained in:
eafonyang
2026-06-12 15:50:01 +08:00
parent 46fd35b0df
commit 8010cbd32f
24 changed files with 2610 additions and 66 deletions
+213
View File
@@ -0,0 +1,213 @@
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
}
+10
View File
@@ -22,6 +22,16 @@ func NewBrandHandler(repo *repository.BrandRepository, log *zap.Logger) *BrandHa
}
}
// GetBrand 获取品牌列表
//
// @Summary 获取品牌列表
// @Tags Brand
// @Produce json
// @Param brandName query string false "品牌名称(模糊匹配)"
// @Param base64Resp query string false "是否返回 base64 编码响应"
// @Success 200 {array} object
// @Failure 500 {object} object
// @Router /audio/getBrand [get]
func (h *BrandHandler) GetBrand(c *gin.Context) {
brandName := c.Query("brandName")
base64Resp := encode.ParseBase64Param(c)
+27
View File
@@ -39,6 +39,19 @@ func NewCurveHandler(repo *repository.CurveRepository, curveCache *cache.CurveCa
}
// ModelCurve 获取机型默认曲线(固定 target: Harman over-ear 2018
//
// @Summary 获取机型默认频响曲线
// @Description 返回指定机型的频响曲线数据(fr),固定使用 Harman over-ear 2018 target
// @Tags Curve
// @Produce json
// @Param brand query string true "品牌名称"
// @Param name query string true "型号名称"
// @Param base64Resp query string false "是否返回 base64 编码响应"
// @Success 200 {object} map[string]any "成功返回 fr 数据"
// @Failure 400 {object} map[string]any "参数校验失败"
// @Failure 500 {object} object
// @Router /audio/modelCurve [get]
//
// GET /audio/modelCurve?brand=xxx&name=xxx&base64Resp=true
func (h *CurveHandler) ModelCurve(c *gin.Context) {
brand := strings.TrimSpace(queryParam(c, "brand"))
@@ -114,6 +127,20 @@ func (h *CurveHandler) ModelCurve(c *gin.Context) {
}
// GetCurve 获取目标曲线
//
// @Summary 获取目标曲线参数化 EQ
// @Description 根据机型和目标曲线名称,计算并返回 parametric_eq 数据
// @Tags Curve
// @Produce json
// @Param brand query string true "品牌名称"
// @Param name query string true "型号名称"
// @Param target query string true "目标曲线名称"
// @Param base64Resp query string false "是否返回 base64 编码响应"
// @Success 200 {object} map[string]any "成功返回 parametric_eq 数据"
// @Failure 400 {object} map[string]any "参数校验失败"
// @Failure 500 {object} object
// @Router /audio/getCurve [get]
//
// GET /audio/getCurve?brand=xxx&name=xxx&target=xxx&base64Resp=true
func (h *CurveHandler) GetCurve(c *gin.Context) {
brand := strings.TrimSpace(queryParam(c, "brand"))
+11
View File
@@ -24,6 +24,17 @@ func NewDeviceHandler(redis *redis.Client, log *zap.Logger) *DeviceHandler {
}
}
// ReportDevInfo 上报设备信息
//
// @Summary 上报设备信息
// @Tags Device
// @Produce json
// @Param mac query string true "设备 MAC 地址"
// @Param model query string true "设备型号"
// @Param ver query string false "固件版本号"
// @Success 200 {object} map[string]any "操作成功"
// @Failure 500 {object} object
// @Router /audio/reportDevInfo [get]
func (h *DeviceHandler) ReportDevInfo(c *gin.Context) {
mac := strings.TrimSpace(c.Query("mac"))
model := strings.TrimSpace(c.Query("model"))
+7
View File
@@ -11,6 +11,13 @@ func NewHealthHandler() *HealthHandler {
return &HealthHandler{}
}
// Check 健康检查
//
// @Summary 健康检查
// @Tags System
// @Produce json
// @Success 200 {object} object
// @Router /api/v1/health [get]
func (h *HealthHandler) Check(c *gin.Context) {
response.OK(c, gin.H{
"status": "up",
+11
View File
@@ -22,6 +22,17 @@ func NewModelHandler(repo *repository.ModelRepository, log *zap.Logger) *ModelHa
}
}
// GetModel 获取型号列表
//
// @Summary 获取型号列表
// @Tags Model
// @Produce json
// @Param brandName query string false "品牌名称"
// @Param modelName query string false "型号名称"
// @Param base64Resp query string false "是否返回 base64 编码响应"
// @Success 200 {array} object
// @Failure 500 {object} object
// @Router /audio/getModel [get]
func (h *ModelHandler) GetModel(c *gin.Context) {
brandName := c.Query("brandName")
modelName := c.Query("modelName")
+14
View File
@@ -28,6 +28,20 @@ func NewModelCSVHandler(s3 *storage.S3Storage, log *zap.Logger) *ModelCSVHandler
}
// GetModelCSV 从 S3 读取耳机 CSV 频响数据
//
// @Summary 获取耳机原始频响 CSV 数据
// @Description 从 S3 读取指定机型的测量 CSV 数据,返回 frequency 和 raw 数组
// @Tags Curve
// @Produce json
// @Param brand query string true "品牌名称"
// @Param model query string true "型号名称"
// @Param form query string true "耳机类型 (in-ear/over-ear)"
// @Param base64Resp query string false "是否返回 base64 编码响应"
// @Success 200 {object} map[string]any "成功返回 frequency 和 raw 数组"
// @Failure 400 {object} map[string]any "参数校验失败"
// @Failure 500 {object} object
// @Router /audio/getModelCSV [get]
//
// GET /audio/getModelCSV?brand=Abyss&model=Dinan DZ&form=over-ear&base64Resp=true
func (h *ModelCSVHandler) GetModelCSV(c *gin.Context) {
brand := strings.TrimSpace(queryParam(c, "brand"))
+11
View File
@@ -23,6 +23,17 @@ func NewModelListHandler(searchClient *search.Client, log *zap.Logger) *ModelLis
}
}
// ModelList 搜索型号列表
//
// @Summary 搜索型号列表(基于 Meilisearch
// @Tags Model
// @Produce json
// @Param key query string false "搜索关键词"
// @Param count query int false "返回数量上限"default(100)
// @Param base64Resp query string false "是否返回 base64 编码响应"
// @Success 200 {array} string
// @Failure 500 {object} object
// @Router /audio/modelList [get]
func (h *ModelListHandler) ModelList(c *gin.Context) {
key := c.Query("key")
base64Resp := encode.ParseBase64Param(c)
+13 -1
View File
@@ -21,7 +21,19 @@ func NewOTAHandler(repo *repository.OTARepository, log *zap.Logger) *OTAHandler
}
// GetOTA 获取 OTA 升级信息
// GET /audio/ota?model=xxx&hw=1&mac=xx:xx:xx&beta=0
//
// @Summary 获取 OTA 升级信息
// @Description 根据设备型号和硬件版本查询最新 OTA 记录,支持黑名单过滤和定向升级逻辑
// @Tags OTA
// @Produce json
// @Param model query string true "设备型号"
// @Param hw query int true "硬件版本"
// @Param mac query string false "设备 MAC 地址"
// @Param beta query int false "是否 beta 通道 (0=否,1=是)"
// @Success 200 {object} object "成功返回 OTA 信息"
// @Failure 400 {object} object
// @Failure 500 {object} object
// @Router /audio/ota [get]
func (h *OTAHandler) GetOTA(c *gin.Context) {
modelVal := strings.TrimSpace(c.Query("model"))
hwStr := strings.TrimSpace(c.Query("hw"))
+170
View File
@@ -0,0 +1,170 @@
package handler
import (
"encoding/json"
"net/http"
"strings"
"github.com/gin-gonic/gin"
"github.com/luxsin/app-api/internal/cache"
"github.com/luxsin/app-api/pkg/encode"
"go.uber.org/zap"
)
type ShareCodeHandler struct {
cache *cache.ShareCodeCache
log *zap.Logger
}
func NewShareCodeHandler(shareCache *cache.ShareCodeCache, log *zap.Logger) *ShareCodeHandler {
return &ShareCodeHandler{cache: shareCache, log: log}
}
// ExportShareCode 导出分享码
//
// @Summary 创建 EQ 分享码
// @Description 将用户的 EQ 数据生成一个 5 位分享码,有效期 30 分钟
// @Tags ShareCode
// @Accept json
// @Produce json
// @Param body body object{mac=string,eq_data=object} true "分享请求"
// @Success 200 {object} map[string]any "成功返回 share_code、expire_at、eq_data"
// @Failure 400 {object} map[string]any "参数校验失败"
// @Failure 500 {object} map[string]any "系统错误"
// @Router /audio/shareCreate [post]
//
// POST /audio/shareCreate
// Body: {"mac": "xx", "eq_data": {...}}
func (h *ShareCodeHandler) ExportShareCode(c *gin.Context) {
var req struct {
Mac string `json:"mac"`
EqData json.RawMessage `json:"eq_data"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusOK, gin.H{
"code": 400,
"msg": "参数校验失败",
})
return
}
mac := strings.TrimSpace(req.Mac)
eqDataRaw := strings.TrimSpace(string(req.EqData))
clientIP := encode.ClientPublicIP(c)
if mac == "" || eqDataRaw == "" {
c.JSON(http.StatusOK, gin.H{
"code": 400,
"msg": "参数校验失败",
})
return
}
if !json.Valid([]byte(eqDataRaw)) {
c.JSON(http.StatusOK, gin.H{
"code": 400,
"msg": "eq_data 格式错误",
})
return
}
ctx := c.Request.Context()
data, err := h.cache.Create(ctx, mac, clientIP, []byte(eqDataRaw))
if err != nil {
h.log.Error("create share code failed", zap.Error(err))
c.JSON(http.StatusOK, gin.H{
"code": 500,
"msg": "系统错误",
})
return
}
var eqData any
if err := json.Unmarshal([]byte(data.EqData), &eqData); err != nil {
eqData = data.EqData
}
c.JSON(http.StatusOK, gin.H{
"code": 200,
"msg": "操作成功",
"share_code": data.ShareCode,
"expire_at": data.ExpireAt.Format("2006-01-02 15:04:05"),
"eq_data": eqData,
})
}
// ImportShareCode 导入分享码
//
// @Summary 导入 EQ 分享码
// @Description 根据分享码获取他人分享的 EQ 数据
// @Tags ShareCode
// @Produce json
// @Param mac query string true "设备 MAC 地址"
// @Param shareCode query string true "5 位分享码"
// @Success 200 {object} map[string]any "成功返回 eq_data"
// @Failure 400 {object} map[string]any "参数校验失败"
// @Failure 500 {object} map[string]any "系统错误"
// @Router /audio/shareAccept [get]
//
// GET /audio/shareAccept?mac=xx&shareCode=ABC12
func (h *ShareCodeHandler) ImportShareCode(c *gin.Context) {
mac := strings.TrimSpace(c.Query("mac"))
shareCode := strings.TrimSpace(c.Query("shareCode"))
clientIP := encode.ClientPublicIP(c)
if mac == "" || shareCode == "" {
c.JSON(http.StatusOK, gin.H{
"code": 400,
"msg": "参数校验失败",
})
return
}
if len(shareCode) != 5 {
c.JSON(http.StatusOK, gin.H{
"code": 0,
"msg": "分享码无效或已过期",
})
return
}
ctx := c.Request.Context()
data, err := h.cache.Get(ctx, shareCode)
if err != nil {
h.log.Error("get share code failed", zap.String("share_code", shareCode), zap.Error(err))
c.JSON(http.StatusOK, gin.H{
"code": 500,
"msg": "系统错误",
})
return
}
if data == nil {
c.JSON(http.StatusOK, gin.H{
"code": 0,
"msg": "分享码无效或已过期",
})
return
}
if err := h.cache.EnqueueImportLog(ctx, mac, shareCode, clientIP, data.EqData, data.ExpireAt); err != nil {
h.log.Error("enqueue share import log failed",
zap.String("share_code", shareCode),
zap.Error(err),
)
c.JSON(http.StatusOK, gin.H{
"code": 500,
"msg": "系统错误",
})
return
}
var eqData any
if err := json.Unmarshal([]byte(data.EqData), &eqData); err != nil {
eqData = data.EqData
}
c.JSON(http.StatusOK, gin.H{
"code": 200,
"msg": "操作成功",
"eq_data": eqData,
})
}
+14
View File
@@ -0,0 +1,14 @@
package model
import "time"
type ShareCodeLog struct {
ID int
MacAddr string
ShareCode string
Action string // export | import
IpAddr string
EqData []byte
ExpireAt *time.Time
CreateAt time.Time
}
+58
View File
@@ -0,0 +1,58 @@
package repository
import (
"context"
"database/sql"
"fmt"
"time"
"github.com/luxsin/app-api/internal/model"
)
type ShareCodeRepository struct {
db *sql.DB
}
func NewShareCodeRepository(db *sql.DB) *ShareCodeRepository {
return &ShareCodeRepository{db: db}
}
func (r *ShareCodeRepository) InsertLog(ctx context.Context, log model.ShareCodeLog) error {
const query = `INSERT INTO share_code_log (mac_addr, share_code, action, ip_addr, eq_data, expire_at)
VALUES (?, ?, ?, ?, ?, ?)`
var expireAt any
if log.ExpireAt != nil {
expireAt = *log.ExpireAt
}
_, err := r.db.ExecContext(ctx, query,
log.MacAddr, log.ShareCode, log.Action, log.IpAddr, log.EqData, expireAt,
)
if err != nil {
return fmt.Errorf("insert share_code_log: %w", err)
}
return nil
}
func (r *ShareCodeRepository) HasExportLog(ctx context.Context, shareCode string) (bool, error) {
const query = `SELECT 1 FROM share_code_log WHERE share_code = ? AND action = 'export' LIMIT 1`
var one int
err := r.db.QueryRowContext(ctx, query, shareCode).Scan(&one)
if err == sql.ErrNoRows {
return false, nil
}
if err != nil {
return false, fmt.Errorf("query share export log: %w", err)
}
return true, nil
}
func ParseShareExpireAt(expireAt time.Time) *time.Time {
if expireAt.IsZero() {
return nil
}
t := expireAt
return &t
}
+11
View File
@@ -13,6 +13,9 @@ import (
"github.com/luxsin/app-api/internal/storage"
"github.com/redis/go-redis/v9"
"go.uber.org/zap"
swaggerFiles "github.com/swaggo/files"
ginSwagger "github.com/swaggo/gin-swagger"
)
func New(log *zap.Logger, db *sql.DB, searchClient *search.Client, rdb *redis.Client, eqCfg config.EqualizeConfig, s3 *storage.S3Storage) *gin.Engine {
@@ -33,6 +36,8 @@ func New(log *zap.Logger, db *sql.DB, searchClient *search.Client, rdb *redis.Cl
otaRepo := repository.NewOTARepository(db)
curveRepo := repository.NewCurveRepository(db)
shareCodeCache := cache.NewShareCodeCache(rdb)
// Handler
health := handler.NewHealthHandler()
brand := handler.NewBrandHandler(brandRepo, log)
@@ -42,12 +47,16 @@ func New(log *zap.Logger, db *sql.DB, searchClient *search.Client, rdb *redis.Cl
ota := handler.NewOTAHandler(otaRepo, log)
curve := handler.NewCurveHandler(curveRepo, curveCache, eqCfg, s3, log)
modelCSV := handler.NewModelCSVHandler(s3, log)
shareCode := handler.NewShareCodeHandler(shareCodeCache, log)
v1 := r.Group("/api/v1")
{
v1.GET("/health", health.Check)
}
// Swagger
r.GET("/docs/*any", ginSwagger.WrapHandler(swaggerFiles.Handler))
audio := r.Group("/audio")
{
audio.GET("/getBrand", brand.GetBrand)
@@ -58,6 +67,8 @@ func New(log *zap.Logger, db *sql.DB, searchClient *search.Client, rdb *redis.Cl
audio.GET("/getCurve", curve.GetCurve)
audio.GET("/modelCurve", curve.ModelCurve)
audio.GET("/getModelCSV", modelCSV.GetModelCSV)
audio.POST("/shareCreate", shareCode.ExportShareCode)
audio.GET("/shareAccept", shareCode.ImportShareCode)
}
return r
+229
View File
@@ -0,0 +1,229 @@
package task
import (
"context"
"encoding/json"
"time"
"github.com/luxsin/app-api/internal/cache"
"github.com/luxsin/app-api/internal/model"
"github.com/luxsin/app-api/internal/repository"
"go.uber.org/zap"
)
type ShareCodePersistTask struct {
cache *cache.ShareCodeCache
repo *repository.ShareCodeRepository
log *zap.Logger
}
func NewShareCodePersistTask(shareCache *cache.ShareCodeCache, repo *repository.ShareCodeRepository, log *zap.Logger) *ShareCodePersistTask {
return &ShareCodePersistTask{cache: shareCache, repo: repo, log: log}
}
func (t *ShareCodePersistTask) Start(interval time.Duration) {
go func() {
ticker := time.NewTicker(interval)
defer ticker.Stop()
for range ticker.C {
t.Persist()
}
}()
t.log.Info("share code persist task started", zap.String("interval", interval.String()))
}
func (t *ShareCodePersistTask) Persist() {
t.persistExports()
t.persistImports()
}
func (t *ShareCodePersistTask) persistExports() {
ctx := context.Background()
codes, err := t.cache.ListPending(ctx)
if err != nil {
t.log.Error("list pending share codes failed", zap.Error(err))
return
}
if len(codes) == 0 {
return
}
t.log.Info("persisting share code export logs", zap.Int("count", len(codes)))
var persisted, skipped, failed int
for _, code := range codes {
ok := t.persistExportOne(ctx, code)
switch ok {
case 1:
persisted++
case 0:
skipped++
default:
failed++
}
}
if persisted > 0 || failed > 0 {
t.log.Info("share code export logs persist finished",
zap.Int("persisted", persisted),
zap.Int("skipped", skipped),
zap.Int("failed", failed),
)
}
}
func (t *ShareCodePersistTask) persistImports() {
ctx := context.Background()
logs, err := t.cache.ListPendingImports(ctx)
if err != nil {
t.log.Error("list pending share import logs failed", zap.Error(err))
return
}
if len(logs) == 0 {
return
}
t.log.Info("persisting share code import logs", zap.Int("count", len(logs)))
var persisted, skipped, failed int
for field, payload := range logs {
ok := t.persistImportOne(ctx, field, payload)
switch ok {
case 1:
persisted++
case 0:
skipped++
default:
failed++
}
}
if persisted > 0 || failed > 0 {
t.log.Info("share code import logs persist finished",
zap.Int("persisted", persisted),
zap.Int("skipped", skipped),
zap.Int("failed", failed),
)
}
}
// persistExportOne 返回值:1=成功刷入,0=跳过,-1=失败
func (t *ShareCodePersistTask) persistExportOne(ctx context.Context, shareCode string) int {
acquired, err := t.cache.AcquireFlushLock(ctx, shareCode)
if err != nil {
t.log.Error("acquire share flush lock failed", zap.String("share_code", shareCode), zap.Error(err))
return -1
}
if !acquired {
return 0
}
defer func() {
if err := t.cache.ReleaseFlushLock(ctx, shareCode); err != nil {
t.log.Warn("release share flush lock failed", zap.String("share_code", shareCode), zap.Error(err))
}
}()
data, err := t.cache.Get(ctx, shareCode)
if err != nil {
t.log.Error("get share code from redis failed", zap.String("share_code", shareCode), zap.Error(err))
return -1
}
if data == nil {
_ = t.cache.MarkPersisted(ctx, shareCode)
return 0
}
if data.Persisted {
return 0
}
exists, err := t.repo.HasExportLog(ctx, shareCode)
if err != nil {
t.log.Error("check share export log failed", zap.String("share_code", shareCode), zap.Error(err))
return -1
}
if exists {
if err := t.cache.MarkPersisted(ctx, shareCode); err != nil {
t.log.Error("mark share code persisted failed", zap.String("share_code", shareCode), zap.Error(err))
return -1
}
return 0
}
logEntry := model.ShareCodeLog{
MacAddr: data.MacAddr,
ShareCode: data.ShareCode,
Action: "export",
IpAddr: data.IpAddr,
EqData: []byte(data.EqData),
ExpireAt: repository.ParseShareExpireAt(data.ExpireAt),
}
if err := t.repo.InsertLog(ctx, logEntry); err != nil {
t.log.Error("insert share export log failed", zap.String("share_code", shareCode), zap.Error(err))
return -1
}
if err := t.cache.MarkPersisted(ctx, shareCode); err != nil {
t.log.Error("mark share code persisted failed", zap.String("share_code", shareCode), zap.Error(err))
return -1
}
return 1
}
// persistImportOne 返回值:1=成功刷入,0=跳过,-1=失败
func (t *ShareCodePersistTask) persistImportOne(ctx context.Context, field, payload string) int {
acquired, err := t.cache.AcquireImportFlushLock(ctx, field)
if err != nil {
t.log.Error("acquire share import flush lock failed", zap.String("field", field), zap.Error(err))
return -1
}
if !acquired {
return 0
}
defer func() {
if err := t.cache.ReleaseImportFlushLock(ctx, field); err != nil {
t.log.Warn("release share import flush lock failed", zap.String("field", field), zap.Error(err))
}
}()
var pending cache.ShareImportPendingLog
if err := json.Unmarshal([]byte(payload), &pending); err != nil {
t.log.Error("unmarshal share import pending log failed", zap.String("field", field), zap.Error(err))
return -1
}
expireAt, err := time.Parse(time.RFC3339, pending.ExpireAt)
if err != nil {
t.log.Error("parse share import expire_at failed", zap.String("field", field), zap.Error(err))
return -1
}
logEntry := model.ShareCodeLog{
MacAddr: pending.MacAddr,
ShareCode: pending.ShareCode,
Action: "import",
IpAddr: pending.IpAddr,
EqData: []byte(pending.EqData),
ExpireAt: repository.ParseShareExpireAt(expireAt),
}
if err := t.repo.InsertLog(ctx, logEntry); err != nil {
t.log.Error("insert share import log failed",
zap.String("share_code", pending.ShareCode),
zap.String("field", field),
zap.Error(err),
)
return -1
}
if err := t.cache.RemovePendingImport(ctx, field); err != nil {
t.log.Error("remove pending share import log failed", zap.String("field", field), zap.Error(err))
return -1
}
return 1
}