新增分享码功能
This commit is contained in:
Vendored
+213
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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
@@ -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"))
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user