Files
app-api/internal/handler/curve.go
T
eafonyang 99785672a5 优化
2026-06-05 20:15:30 +08:00

494 lines
13 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/csv"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"os"
"strconv"
"strings"
"time"
"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/pkg/encode"
"go.uber.org/zap"
)
type CurveHandler struct {
repo *repository.CurveRepository
cache *cache.CurveCache
cfg config.EqualizeConfig
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) {
brand := strings.TrimSpace(queryParam(c, "brand"))
name := strings.TrimSpace(queryParam(c, "name"))
target := strings.TrimSpace(queryParam(c, "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 接口并缓存结果
func (h *CurveHandler) getCurvePoint(ctx context.Context, brand, name, target string) (string, error) {
data, acquired, err := h.cache.GetWithLock(ctx, brand, name, target)
if err != nil {
return "", err
}
// 缓存命中
if data != "" && !acquired {
h.log.Info("curve cache hit", zap.String("brand", brand), zap.String("name", name), zap.String("target", target))
return data, nil
}
// 获取了锁,缓存仍然为空,需要请求 EQ 接口
if acquired {
h.log.Info("curve cache miss, requesting eq api", zap.String("brand", brand), zap.String("name", name), zap.String("target", target))
result, eqErr := h.getCurvePointFromPEQ(ctx, brand, name, target)
if eqErr != nil {
// 释放锁
if lockErr := h.cache.ReleaseLock(ctx, brand, name, target); lockErr != nil {
h.log.Warn("release curve lock failed", zap.Error(lockErr))
}
return "", eqErr
}
// 释放锁
if lockErr := h.cache.ReleaseLock(ctx, brand, name, target); lockErr != nil {
h.log.Warn("release curve lock failed", zap.Error(lockErr))
}
if result != "" {
// 缓存结果
if cacheErr := h.cache.Set(ctx, brand, name, target, result); cacheErr != nil {
h.log.Warn("curve cache set failed", zap.Error(cacheErr))
}
return result, nil
}
return "", nil
}
// 未获取锁(其他请求正在处理),等待后重试
time.Sleep(100 * time.Millisecond)
return h.getCurvePoint(ctx, brand, name, target)
}
// 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 源需要读 CSV
var measurement map[string]any
if m.Source != nil && *m.Source == "Eafonyoung" {
measurement, err = h.readCSV(h.cfg.MeasurementBasePath + "/Eafonyoung/data/" + deref(m.Form) + "/" + headPhone + ".csv")
if err != nil || measurement == nil {
return "", nil
}
}
// 构建请求
var targetParam any
var targetRaw map[string]any
if bool(t.ReadCSV) {
targetRaw, err = h.readCSV(h.cfg.TargetBasePath + "/" + deref(t.File))
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}
}
resp, err := h.reqEqualize(headPhone, measurement, targetParam, bassBoost, deref(m.Source), deref(m.Rig))
if err != nil {
return "", err
}
return resp, nil
}
// reqEqualize 调用 autoeq API
func (h *CurveHandler) reqEqualize(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, h.cfg.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", h.cfg.APIURL),
)
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
}
// readCSV 读取 CSV 文件并返回 {frequency: [...], raw: [...]}
func (h *CurveHandler) readCSV(path string) (map[string]any, error) {
f, err := os.Open(path)
if err != nil {
if os.IsNotExist(err) {
h.log.Info("csv file not found", zap.String("path", path))
return nil, nil
}
return nil, fmt.Errorf("open csv: %w", err)
}
defer f.Close()
reader := csv.NewReader(f)
// 跳过表头
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 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
}
}
// queryParam 从 URL 原始 query 中获取参数,保留 + 为字面量而非空格
func queryParam(c *gin.Context, key string) string {
vals, ok := c.Request.URL.Query()[key]
if !ok || len(vals) == 0 {
return ""
}
// c.Query() 会把 + 解码为空格,这里从原始 query 手动解码,+ 保留为 +
if strings.Contains(vals[0], " ") {
rawQuery := c.Request.URL.RawQuery
for _, pair := range strings.Split(rawQuery, "&") {
kv := strings.SplitN(pair, "=", 2)
if len(kv) == 2 && kv[0] == key {
decoded, err := url.PathUnescape(strings.ReplaceAll(kv[1], "+", "%2B"))
if err == nil {
return decoded
}
}
}
}
return vals[0]
}