Files
app-api/internal/handler/curve.go
T

558 lines
16 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package handler
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"
"github.com/aws/aws-sdk-go-v2/service/s3/types"
"github.com/gin-gonic/gin"
"github.com/luxsin/app-api/internal/cache"
"github.com/luxsin/app-api/internal/config"
"github.com/luxsin/app-api/internal/repository"
"github.com/luxsin/app-api/internal/response"
"github.com/luxsin/app-api/internal/storage"
"github.com/luxsin/app-api/pkg/encode"
"go.uber.org/zap"
)
type CurveHandler struct {
repo *repository.CurveRepository
cache *cache.CurveCache
cfg config.EqualizeConfig
s3 *storage.S3Storage
log *zap.Logger
}
const defaultModelCurveTarget = "Harman over-ear 2018"
func NewCurveHandler(repo *repository.CurveRepository, curveCache *cache.CurveCache, cfg config.EqualizeConfig, s3 *storage.S3Storage, log *zap.Logger) *CurveHandler {
return &CurveHandler{repo: repo, cache: curveCache, cfg: cfg, s3: s3, log: log}
}
// 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(c.Query("brand"))
name := strings.TrimSpace(c.Query("name"))
base64Resp := encode.ParseBase64Param(c)
if brand == "" || name == "" {
h.writeCurveResponse(c, base64Resp, gin.H{
"code": 400,
"msg": "参数校验失败",
})
return
}
ctx := c.Request.Context()
// 快速路径:直接读取独立存储的 fr 数据(缓存命中时只需一次 Redis 调用)
frData, err := h.cache.GetFR(ctx, brand, name)
if err != nil {
h.log.Warn("get fr from cache failed", zap.Error(err))
}
if frData != "" {
h.log.Info("model curve fr cache hit", zap.String("brand", brand), zap.String("name", name))
var fr any
if err := json.Unmarshal([]byte(frData), &fr); err == nil {
h.writeCurveResponse(c, base64Resp, gin.H{
"code": 200,
"msg": "ok",
"fr": fr,
})
return
}
}
// 慢速路径:缓存未命中,走完整流程获取数据(同时会写入 __fr 缓存)
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
}
// 只返回 fr,不包含 parametric_eq
fr, hasFR := resp["fr"]
if !hasFR {
h.writeCurveResponse(c, base64Resp, gin.H{
"code": 0,
"msg": "无曲线数据",
})
return
}
h.writeCurveResponse(c, base64Resp, gin.H{
"code": 200,
"msg": "ok",
"fr": fr,
})
}
// 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(c.Query("brand"))
name := strings.TrimSpace(c.Query("name"))
target := strings.TrimSpace(c.Query("target"))
base64Resp := encode.ParseBase64Param(c)
if brand == "" || name == "" || target == "" {
response.BadRequest(c, "brand、name和target参数不能为空")
return
}
ctx := c.Request.Context()
result, err := h.getCurvePoint(ctx, brand, name, target)
if err != nil {
h.log.Error("get curve point failed", zap.Error(err))
response.InternalError(c, "获取曲线数据失败")
return
}
if result == "" {
response.Fail(c, http.StatusOK, 0, "无曲线数据")
return
}
// 解析 JSON 结果,只提取 parametric_eq
var resp map[string]any
if err := json.Unmarshal([]byte(result), &resp); err != nil {
response.InternalError(c, "解析曲线数据失败")
return
}
parametricEq, _ := resp["parametric_eq"]
resultData := gin.H{
"code": 200,
"msg": "操作成功",
"parametric_eq": parametricEq,
}
if base64Resp {
encoded, err := encode.EncodeJSON(resultData)
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, 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 接口并缓存结果
// 缓存结构优化:fr 数据(与 target 无关)单独存储在 __fr 字段,避免每个 target 重复存储
func (h *CurveHandler) getCurvePoint(ctx context.Context, brand, name, target string) (string, error) {
// 使用 GetWithFR,命中缓存时自动合并 __fr 数据
data, err := h.cache.GetWithFR(ctx, brand, name, target)
if err != nil {
return "", err
}
if data != "" {
h.log.Info("curve cache hit", zap.String("brand", brand), zap.String("name", name), zap.String("target", target))
return data, nil
}
// 缓存未命中,获取分布式锁
acquired, lockErr := h.cache.AcquireLock(ctx, brand, name, target)
if lockErr != nil {
return "", lockErr
}
// 未获取锁(其他请求正在处理),等待后重试
if !acquired {
time.Sleep(100 * time.Millisecond)
return h.getCurvePoint(ctx, brand, name, target)
}
// 获取了锁,双重检查缓存
h.log.Info("curve cache miss, requesting eq api", zap.String("brand", brand), zap.String("name", name), zap.String("target", target))
data, err = h.cache.GetWithFR(ctx, brand, name, target)
if err != nil {
h.releaseLock(brand, name, target)
return "", err
}
if data != "" {
h.releaseLock(brand, name, target)
return data, nil
}
// 请求 EQ 接口
result, eqErr := h.getCurvePointFromPEQ(ctx, brand, name, target)
if eqErr != nil {
h.releaseLock(brand, name, target)
return "", eqErr
}
if result == "" {
h.releaseLock(brand, name, target)
return "", nil
}
// 将 fr 数据剥离,单独存储在 __fr 字段,避免每个 target 重复存储
dataToCache := result
var frToCache string
var parsed map[string]json.RawMessage
if err := json.Unmarshal([]byte(result), &parsed); err == nil {
if fr, ok := parsed["fr"]; ok {
frToCache = string(fr)
delete(parsed, "fr")
if stripped, err := json.Marshal(parsed); err == nil {
dataToCache = string(stripped)
}
}
}
// 先写缓存,再释放锁(避免其他服务器在缓存写入前抢到锁后重复调用 EQ API)
if cacheErr := h.cache.Set(ctx, brand, name, target, dataToCache); cacheErr != nil {
h.log.Warn("curve cache set failed", zap.Error(cacheErr))
}
if frToCache != "" {
existingFR, _ := h.cache.Get(ctx, brand, name, "__fr")
if existingFR == "" {
if cacheErr := h.cache.Set(ctx, brand, name, "__fr", frToCache); cacheErr != nil {
h.log.Warn("curve fr cache set failed", zap.Error(cacheErr))
}
}
}
h.releaseLock(brand, name, target)
return result, nil
}
// releaseLock 释放分布式锁并记录警告日志
func (h *CurveHandler) releaseLock(brand, name, target string) {
if err := h.cache.ReleaseLock(context.Background(), brand, name, target); err != nil {
h.log.Warn("release curve lock failed", zap.Error(err))
}
}
// getCurvePointFromPEQ 从 EQ 接口获取曲线数据
func (h *CurveHandler) getCurvePointFromPEQ(ctx context.Context, brand, name, targetName string) (string, error) {
m, err := h.repo.GetModelByBrandAndName(ctx, brand, name)
if err != nil {
return "", fmt.Errorf("query model: %w", err)
}
if m == nil {
return "", nil
}
t, err := h.repo.GetTargetByLabel(ctx, targetName)
if err != nil {
return "", fmt.Errorf("query target: %w", err)
}
if t == nil {
return "", nil
}
// 确定 headPhone 名称
headPhone := brand + " " + name
if m.EqKey != nil && *m.EqKey == "name" {
headPhone = name
}
// 获取 measurement 数据(仅 Eafonyoung 源需要从 S3 读 CSV
var measurement map[string]any
if m.Source != nil && *m.Source == "Eafonyoung" {
key := modelCSVKey(brand, name, deref(m.Form))
measurement, err = h.readCSVFromS3(ctx, key)
if err != nil || measurement == nil {
return "", nil
}
}
// 构建请求
var targetParam any
var targetRaw map[string]any
if bool(t.ReadCSV) {
key := targetCSVKey(deref(t.File))
targetRaw, err = h.readCSVFromS3(ctx, key)
if err != nil || targetRaw == nil {
return "", nil
}
targetParam = targetRaw
} else {
targetParam = targetName
}
// 解析 bassBoost
var bassBoost map[string]any
if t.BassBoost != nil {
if err := json.Unmarshal([]byte(*t.BassBoost), &bassBoost); err != nil {
h.log.Warn("parse bassBoost failed", zap.Error(err))
bassBoost = map[string]any{"gain": 0, "fc": 100, "q": 0.7}
}
} else {
bassBoost = map[string]any{"gain": 0, "fc": 100, "q": 0.7}
}
apiURL := h.eqAPIURL(deref(m.Source))
resp, err := h.reqEqualize(apiURL, headPhone, measurement, targetParam, bassBoost, deref(m.Source), deref(m.Rig))
if err != nil {
return "", err
}
return resp, nil
}
// eqAPIURL 正式环境非 Eafonyoung 源走公网 EQ,其余走内网/默认地址
func (h *CurveHandler) eqAPIURL(source string) string {
if source != "Eafonyoung" && h.cfg.ExternalAPIURL != "" && h.cfg.ExternalAPIURL != h.cfg.APIURL {
return h.cfg.ExternalAPIURL
}
return h.cfg.APIURL
}
// reqEqualize 调用 autoeq API
func (h *CurveHandler) reqEqualize(apiURL, headPhone string, measurement map[string]any, target any, bassBoost map[string]any, source, rig string) (string, error) {
gain := floatVal(bassBoost, "gain", 0)
fc := intVal(bassBoost, "fc", 100)
q := floatVal(bassBoost, "q", 0.7)
reqBody := map[string]any{
"target": target,
"sound_signature": nil,
"sound_signature_smoothing_window_size": 1,
"bass_boost_gain": gain,
"bass_boost_fc": fc,
"bass_boost_q": q,
"treble_boost_gain": 0,
"treble_boost_fc": 10000,
"treble_boost_q": 0.7,
"tilt": 0,
"fs": 48000,
"bit_depth": 16,
"phase": "minimum",
"f_res": 16,
"preamp": 0,
"max_gain": 12,
"max_slope": 18,
"window_size": 0.08,
"treble_window_size": 2,
"treble_f_lower": 6000,
"treble_f_upper": 8000,
"treble_gain_k": 1,
"graphic_eq": false,
"parametric_eq": true,
"fixed_band_eq": false,
"convolution_eq": false,
"source": source,
"rig": rig,
"parametric_eq_config": "MINIDSP_IL_DSP",
"response": map[string]any{
"fr_f_step": 1.02,
"base64fp16": false,
"fr_fields": []string{"raw"},
},
}
// measurement 或 headPhone 二选一
if measurement != nil {
reqBody["measurement"] = measurement
} else if headPhone != "" {
reqBody["name"] = headPhone
}
jsonData, err := json.Marshal(reqBody)
if err != nil {
return "", fmt.Errorf("marshal eq request: %w", err)
}
httpReq, err := http.NewRequestWithContext(context.Background(), http.MethodPost, apiURL, bytes.NewReader(jsonData))
if err != nil {
return "", fmt.Errorf("create eq request: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
start := time.Now()
client := &http.Client{Timeout: 10 * time.Second}
httpResp, err := client.Do(httpReq)
elapsed := time.Since(start)
h.log.Info("eq api response",
zap.String("headPhone", headPhone),
zap.Duration("latency", elapsed),
zap.String("url", apiURL),
zap.String("source", source),
)
if err != nil {
h.log.Error("eq api request failed", zap.Duration("latency", elapsed), zap.Error(err))
return "", fmt.Errorf("call eq api: %w", err)
}
defer httpResp.Body.Close()
body, err := io.ReadAll(httpResp.Body)
if err != nil {
return "", fmt.Errorf("read eq response: %w", err)
}
if httpResp.StatusCode != http.StatusOK {
h.log.Warn("eq api returned non-200", zap.Int("status", httpResp.StatusCode), zap.String("body", string(body)))
return "", fmt.Errorf("eq api status: %d", httpResp.StatusCode)
}
var eqResp map[string]any
if err := json.Unmarshal(body, &eqResp); err != nil {
return "", fmt.Errorf("unmarshal eq response: %w", err)
}
parametricEq, _ := eqResp["parametric_eq"].(map[string]any)
if parametricEq == nil {
result := map[string]any{
"code": 200,
"msg": "req param error",
"param": reqBody,
}
resultJSON, _ := json.Marshal(result)
return string(resultJSON), nil
}
result := map[string]any{
"code": 200,
"msg": "ok",
"parametric_eq": eqResp["parametric_eq"],
"fr": eqResp["fr"],
}
resultJSON, _ := json.Marshal(result)
return string(resultJSON), nil
}
// readCSVFromS3 从 S3 读取 CSV 并返回 {frequency: [...], raw: [...]}
func (h *CurveHandler) readCSVFromS3(ctx context.Context, key string) (map[string]any, error) {
h.log.Info("reading csv from s3", zap.String("key", key))
data, err := h.s3.GetObject(ctx, key)
if err != nil {
var noSuchKey *types.NoSuchKey
if errors.As(err, &noSuchKey) {
h.log.Info("csv not found in s3", zap.String("key", key))
return nil, nil
}
return nil, fmt.Errorf("get csv from s3: %w", err)
}
parsed, err := parseCSVData(data)
if err != nil {
return nil, fmt.Errorf("parse csv: %w", err)
}
return parsed, nil
}
// 辅助函数
func deref(s *string) string {
if s == nil {
return ""
}
return *s
}
func floatVal(m map[string]any, key string, defaultVal float64) float64 {
v, ok := m[key]
if !ok {
return defaultVal
}
switch n := v.(type) {
case float64:
return n
case int:
return float64(n)
case string:
f, err := strconv.ParseFloat(n, 64)
if err != nil {
return defaultVal
}
return f
default:
return defaultVal
}
}
func intVal(m map[string]any, key string, defaultVal int) int {
v, ok := m[key]
if !ok {
return defaultVal
}
switch n := v.(type) {
case float64:
return int(n)
case int:
return n
case string:
i, err := strconv.Atoi(n)
if err != nil {
return defaultVal
}
return i
default:
return defaultVal
}
}