This commit is contained in:
eafonyang
2026-06-05 20:15:30 +08:00
parent 8c22be4e0a
commit 99785672a5
15 changed files with 439 additions and 8 deletions
+3
View File
@@ -14,6 +14,7 @@ type Config struct {
Meilisearch MeilisearchConfig
Redis RedisConfig
Equalize EqualizeConfig
S3 S3Config
EnablePersistTask bool
}
@@ -44,6 +45,7 @@ func Load() (*Config, error) {
}
eq := loadEqualize(env)
s3cfg := loadS3(env)
return &Config{
Env: env,
@@ -53,6 +55,7 @@ func Load() (*Config, error) {
Meilisearch: ms,
Redis: rd,
Equalize: eq,
S3: s3cfg,
EnablePersistTask: getEnv("ENABLE_PERSIST_TASK", "false") == "true",
}, nil
}
+36
View File
@@ -0,0 +1,36 @@
package config
import "os"
type S3Config struct {
Bucket string
Region string
AccessKeyID string
SecretAccessKey string
}
func loadS3(env string) S3Config {
if os.Getenv("AWS_ACCESS_KEY_ID") != "" || os.Getenv("AWS_SECRET_ACCESS_KEY") != "" {
return S3Config{
Bucket: getEnv("S3_BUCKET", "luxsin-app-bucket"),
Region: getEnv("AWS_REGION", "eu-central-1"),
AccessKeyID: os.Getenv("AWS_ACCESS_KEY_ID"),
SecretAccessKey: os.Getenv("AWS_SECRET_ACCESS_KEY"),
}
}
switch env {
case "production":
return S3Config{
Bucket: getEnv("S3_BUCKET", "luxsin-app-bucket"),
Region: getEnv("AWS_REGION", "eu-central-1"),
}
default:
return S3Config{
Bucket: getEnv("S3_BUCKET", "luxsin-app-bucket"),
Region: getEnv("AWS_REGION", "eu-central-1"),
AccessKeyID: "AKIAVMFK45I6P3TNI3VJ",
SecretAccessKey: "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",
}
}
}
+58 -1
View File
@@ -30,10 +30,53 @@ type CurveHandler struct {
log *zap.Logger
}
const defaultModelCurveTarget = "Harman over-ear 2018"
func NewCurveHandler(repo *repository.CurveRepository, curveCache *cache.CurveCache, cfg config.EqualizeConfig, log *zap.Logger) *CurveHandler {
return &CurveHandler{repo: repo, cache: curveCache, cfg: cfg, log: log}
}
// ModelCurve 获取机型默认曲线(固定 target: Harman over-ear 2018
// GET /audio/modelCurve?brand=xxx&name=xxx&base64Resp=true
func (h *CurveHandler) ModelCurve(c *gin.Context) {
brand := strings.TrimSpace(queryParam(c, "brand"))
name := strings.TrimSpace(queryParam(c, "name"))
base64Resp := encode.ParseBase64Param(c)
if brand == "" || name == "" {
h.writeCurveResponse(c, base64Resp, gin.H{
"code": 400,
"msg": "参数校验失败",
})
return
}
ctx := c.Request.Context()
result, err := h.getCurvePoint(ctx, brand, name, defaultModelCurveTarget)
if err != nil {
h.log.Error("get model curve failed", zap.Error(err))
response.InternalError(c, "获取曲线数据失败")
return
}
if result == "" {
h.writeCurveResponse(c, base64Resp, gin.H{
"code": 0,
"msg": "无曲线数据",
})
return
}
var resp map[string]any
if err := json.Unmarshal([]byte(result), &resp); err != nil {
response.InternalError(c, "解析曲线数据失败")
return
}
h.writeCurveResponse(c, base64Resp, resp)
}
// GetCurve 获取目标曲线
// GET /audio/getCurve?brand=xxx&name=xxx&target=xxx&base64Resp=true
func (h *CurveHandler) GetCurve(c *gin.Context) {
@@ -57,7 +100,7 @@ func (h *CurveHandler) GetCurve(c *gin.Context) {
}
if result == "" {
response.Fail(c, http.StatusOK, 40004, "无曲线数据")
response.Fail(c, http.StatusOK, 0, "无曲线数据")
return
}
@@ -90,6 +133,20 @@ func (h *CurveHandler) GetCurve(c *gin.Context) {
c.JSON(http.StatusOK, resultData)
}
func (h *CurveHandler) writeCurveResponse(c *gin.Context, base64Resp bool, data any) {
if base64Resp {
encoded, err := encode.EncodeJSON(data)
if err != nil {
h.log.Error("encode response failed", zap.Error(err))
response.InternalError(c, "编码响应失败")
return
}
c.String(http.StatusOK, encoded)
return
}
c.JSON(http.StatusOK, data)
}
// getCurvePoint 获取曲线数据:先查缓存,缓存不存在则请求 EQ 接口并缓存结果
func (h *CurveHandler) getCurvePoint(ctx context.Context, brand, name, target string) (string, error) {
data, acquired, err := h.cache.GetWithLock(ctx, brand, name, target)
+161
View File
@@ -0,0 +1,161 @@
package handler
import (
"bytes"
"encoding/csv"
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"unicode"
"github.com/aws/aws-sdk-go-v2/service/s3/types"
"github.com/gin-gonic/gin"
"github.com/luxsin/app-api/internal/storage"
"github.com/luxsin/app-api/pkg/encode"
"go.uber.org/zap"
)
type ModelCSVHandler struct {
s3 *storage.S3Storage
log *zap.Logger
}
func NewModelCSVHandler(s3 *storage.S3Storage, log *zap.Logger) *ModelCSVHandler {
return &ModelCSVHandler{s3: s3, log: log}
}
// GetModelCSV 从 S3 读取耳机 CSV 频响数据
// 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"))
model := strings.TrimSpace(queryParam(c, "model"))
form := strings.TrimSpace(queryParam(c, "form"))
base64Resp := encode.ParseBase64Param(c)
if brand == "" || model == "" || form == "" {
h.writeModelCSVResponse(c, base64Resp, gin.H{
"code": 400,
"msg": "参数校验失败",
})
return
}
key := modelCSVKey(brand, model, form)
ctx := c.Request.Context()
data, err := h.s3.GetObject(ctx, key)
if err != nil {
var noSuchKey *types.NoSuchKey
if errors.As(err, &noSuchKey) {
h.log.Info("model csv not found in s3", zap.String("key", key))
h.writeModelCSVResponse(c, base64Resp, gin.H{
"code": 0,
"msg": "无曲线数据",
})
return
}
h.log.Error("get model csv from s3 failed",
zap.String("key", key),
zap.Error(err),
)
h.writeModelCSVResponse(c, base64Resp, gin.H{
"code": 500,
"msg": "系统错误",
})
return
}
parsed, err := parseCSVData(data)
if err != nil || parsed == nil {
h.log.Error("parse model csv failed", zap.String("key", key), zap.Error(err))
h.writeModelCSVResponse(c, base64Resp, gin.H{
"code": 0,
"msg": "无曲线数据",
})
return
}
h.writeModelCSVResponse(c, base64Resp, gin.H{
"code": 200,
"msg": "操作成功",
"frequency": parsed["frequency"],
"raw": parsed["raw"],
})
}
func modelCSVKey(brand, model, form string) string {
filename := brand + " " + model + ".csv"
return fmt.Sprintf("autoeq/measurements/Eafonyoung/data/%s/%s/%s",
form, brandPrefix(brand), filename)
}
func brandPrefix(brand string) string {
runes := []rune(brand)
if len(runes) == 0 {
return ""
}
first := runes[0]
if unicode.IsLetter(first) && unicode.IsLower(first) {
return strings.ToUpper(string(first))
}
return string(first)
}
func parseCSVData(data []byte) (map[string]any, error) {
reader := csv.NewReader(bytes.NewReader(data))
if _, err := reader.Read(); err != nil {
return nil, fmt.Errorf("read csv header: %w", err)
}
frequency := make([]float64, 0)
raw := make([]float64, 0)
for {
record, err := reader.Read()
if err == io.EOF {
break
}
if err != nil {
return nil, fmt.Errorf("read csv row: %w", err)
}
if len(record) < 2 {
continue
}
freq, err1 := strconv.ParseFloat(record[0], 64)
val, err2 := strconv.ParseFloat(record[1], 64)
if err1 != nil || err2 != nil {
continue
}
frequency = append(frequency, freq)
raw = append(raw, val)
}
if len(frequency) == 0 {
return nil, nil
}
return map[string]any{
"frequency": frequency,
"raw": raw,
}, nil
}
func (h *ModelCSVHandler) writeModelCSVResponse(c *gin.Context, base64Resp bool, data any) {
if base64Resp {
encoded, err := encode.EncodeJSON(data)
if err != nil {
h.log.Error("encode response failed", zap.Error(err))
c.JSON(http.StatusOK, gin.H{
"code": 500,
"msg": "系统错误",
})
return
}
c.String(http.StatusOK, encoded)
return
}
c.JSON(http.StatusOK, data)
}
+5 -1
View File
@@ -10,11 +10,12 @@ import (
"github.com/luxsin/app-api/internal/middleware"
"github.com/luxsin/app-api/internal/repository"
"github.com/luxsin/app-api/internal/search"
"github.com/luxsin/app-api/internal/storage"
"github.com/redis/go-redis/v9"
"go.uber.org/zap"
)
func New(log *zap.Logger, db *sql.DB, searchClient *search.Client, rdb *redis.Client, eqCfg config.EqualizeConfig) *gin.Engine {
func New(log *zap.Logger, db *sql.DB, searchClient *search.Client, rdb *redis.Client, eqCfg config.EqualizeConfig, s3 *storage.S3Storage) *gin.Engine {
r := gin.New()
r.Use(gin.Recovery())
r.Use(middleware.RequestID())
@@ -40,6 +41,7 @@ func New(log *zap.Logger, db *sql.DB, searchClient *search.Client, rdb *redis.Cl
device := handler.NewDeviceHandler(rdb, log)
ota := handler.NewOTAHandler(otaRepo, log)
curve := handler.NewCurveHandler(curveRepo, curveCache, eqCfg, log)
modelCSV := handler.NewModelCSVHandler(s3, log)
v1 := r.Group("/api/v1")
{
@@ -54,6 +56,8 @@ func New(log *zap.Logger, db *sql.DB, searchClient *search.Client, rdb *redis.Cl
audio.GET("/reportDevInfo", device.ReportDevInfo)
audio.GET("/ota", ota.GetOTA)
audio.GET("/getCurve", curve.GetCurve)
audio.GET("/modelCurve", curve.ModelCurve)
audio.GET("/getModelCSV", modelCSV.GetModelCSV)
}
return r
+56
View File
@@ -0,0 +1,56 @@
package storage
import (
"context"
"fmt"
"io"
"github.com/aws/aws-sdk-go-v2/aws"
awsconfig "github.com/aws/aws-sdk-go-v2/config"
"github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/service/s3"
"github.com/luxsin/app-api/internal/config"
)
type S3Storage struct {
client *s3.Client
bucket string
}
func NewS3Storage(ctx context.Context, cfg config.S3Config) (*S3Storage, error) {
var opts []func(*awsconfig.LoadOptions) error
opts = append(opts, awsconfig.WithRegion(cfg.Region))
if cfg.AccessKeyID != "" && cfg.SecretAccessKey != "" {
opts = append(opts, awsconfig.WithCredentialsProvider(
credentials.NewStaticCredentialsProvider(cfg.AccessKeyID, cfg.SecretAccessKey, ""),
))
}
awsCfg, err := awsconfig.LoadDefaultConfig(ctx, opts...)
if err != nil {
return nil, fmt.Errorf("load aws config: %w", err)
}
return &S3Storage{
client: s3.NewFromConfig(awsCfg),
bucket: cfg.Bucket,
}, nil
}
func (s *S3Storage) GetObject(ctx context.Context, key string) ([]byte, error) {
out, err := s.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
})
if err != nil {
return nil, err
}
defer out.Body.Close()
data, err := io.ReadAll(out.Body)
if err != nil {
return nil, fmt.Errorf("read s3 object body: %w", err)
}
return data, nil
}
+22 -4
View File
@@ -80,14 +80,32 @@ func (t *DevicePersistTask) Persist() {
}
}
// 批量删除成功处理的记录
if len(succeeded) > 0 {
total := len(devices)
persisted := len(succeeded)
failed := total - persisted
// 批量删除成功刷入的记录,失败的保留在 Redis 等待下次重试
if persisted > 0 {
if err := t.rdb.HDel(ctx, "devices", succeeded...).Err(); err != nil {
t.log.Error("redis HDel failed", zap.Int("count", len(succeeded)), zap.Error(err))
t.log.Error("redis HDel failed after persist",
zap.Int("persisted", persisted),
zap.Error(err),
)
} else {
t.log.Info("devices persisted and removed from redis", zap.Int("count", len(succeeded)))
t.log.Info("devices persisted to database and removed from redis",
zap.Int("persisted", persisted),
zap.Int("total", total),
zap.Int("failed", failed),
)
}
}
if failed > 0 {
t.log.Warn("some devices failed to persist, kept in redis for retry",
zap.Int("failed", failed),
zap.Int("total", total),
)
}
}
// persistDevice 处理 user_device 表写入