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} } 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 } // 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] }