Files
app-api/internal/handler/curve.go
T
eafonyang 173e23f531 feat(equalize): Enhance EqualizeConfig to support external API URL
- Added ExternalAPIURL to EqualizeConfig for public EQ interface usage.
- Updated loadEqualize function to initialize ExternalAPIURL from environment variable or default value.
- Refactored reqEqualize method to accept apiURL as a parameter for improved flexibility.
- Introduced eqAPIURL method to determine the appropriate API URL based on the source.
2026-06-08 18:48:54 +08:00

478 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/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"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
// 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 使用标准 query 解码:+ 表示空格(如 Harman+over-ear+2018 → Harman over-ear 2018
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 接口并缓存结果
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 源需要从 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
}
}
// 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]
}