diff --git a/adapter/volcano/volcano.go b/adapter/volcano/volcano.go index 603ff7c..d602075 100644 --- a/adapter/volcano/volcano.go +++ b/adapter/volcano/volcano.go @@ -13,7 +13,6 @@ import ( "time" "github.com/google/uuid" - "github.com/volcano-tts/tts-api/common" "github.com/volcano-tts/tts-api/dto" ) @@ -24,7 +23,6 @@ type HTTPClient struct { func NewHTTPClient() *HTTPClient { return &HTTPClient{ client: &http.Client{ - Timeout: common.DefaultTimeout, Transport: &http.Transport{ MaxIdleConns: 100, MaxIdleConnsPerHost: 20, @@ -52,19 +50,25 @@ func (h *HTTPClient) PostStream(url string, headers map[string]string, body []by } func convertSpeedToSpeechRate(speed float64) int { - if speed <= 0.5 { - return -50 + rate := int((speed - 1.0) * 100) + if rate < -200 { + rate = -200 } - if speed >= 2.0 { - return 100 + if rate > 500 { + rate = 500 } - return int((speed - 1.0) * 100) + return rate } -func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text string, speed float64) (*dto.SynthesisResult, error) { +func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text string, speed float64, voice string) (*dto.SynthesisResult, error) { reqID := uuid.NewString() speechRate := convertSpeedToSpeechRate(speed) + speaker := config.Speaker + if voice != "" { + speaker = voice + } + params := map[string]interface{}{ "user": map[string]interface{}{ "uid": "uid", @@ -72,7 +76,7 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri "namespace": "BidirectionalTTS", "req_params": map[string]interface{}{ "text": text, - "speaker": config.Speaker, + "speaker": speaker, "audio_params": map[string]interface{}{ "format": "wav", "sample_rate": 24000, diff --git a/controller/tts.go b/controller/tts.go index 5b5828c..2b9b62f 100644 --- a/controller/tts.go +++ b/controller/tts.go @@ -43,6 +43,7 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) { body, err := io.ReadAll(r.Body) if err != nil { if strings.Contains(err.Error(), "request body too large") { + http.Error(w, "Request body too large", http.StatusRequestEntityTooLarge) return } http.Error(w, "Failed to read request body", http.StatusBadRequest) @@ -88,7 +89,7 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) { } ttsStart := time.Now() - result, err := volcano.Synthesis(&setting.TTSConfig, volcanoClient, req.Input, speed) + result, err := volcano.Synthesis(&setting.TTSConfig, volcanoClient, req.Input, speed, req.Voice) duration := time.Since(ttsStart) if err != nil { diff --git a/main.go b/main.go index bb557fe..c8c6ada 100644 --- a/main.go +++ b/main.go @@ -47,7 +47,7 @@ func main() { server := &http.Server{ Addr: ":" + port, - Handler: r, + Handler: middleware.CORS(r), ReadTimeout: 30 * time.Second, WriteTimeout: 120 * time.Second, IdleTimeout: 60 * time.Second, diff --git a/middleware/auth.go b/middleware/auth.go index f300588..1a715f4 100644 --- a/middleware/auth.go +++ b/middleware/auth.go @@ -20,7 +20,8 @@ func InitAPIKeys() { } log.Printf("宸查厤缃?%d 涓湁鏁堢殑API瀵嗛挜", len(validAPIKeys)) } else { - log.Println("警告: OPENAI_TTS_API_KEY环境变量未设置,将拒绝所有请求?) + log.Println("警告: OPENAI_TTS_API_KEY环境变量未设置,所有请求将无需认证即可访问") + log.Println("如需启用API密钥验证,请设置 OPENAI_TTS_API_KEY 环境变量(多个密钥用逗号分隔)") } } diff --git a/middleware/cors.go b/middleware/cors.go index 616d091..3f3850d 100644 --- a/middleware/cors.go +++ b/middleware/cors.go @@ -79,35 +79,45 @@ func matchOrigin(origin string) (string, bool) { func CORS(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { origin := r.Header.Get("Origin") - isPreflight := r.Method == http.MethodOptions - if origin != "" { - allowOrigin, matched := matchOrigin(origin) - if matched { - w.Header().Set("Access-Control-Allow-Origin", allowOrigin) - w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") - w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization") - w.Header().Set("Access-Control-Expose-Headers", "X-Request-Id") - w.Header().Set("Access-Control-Max-Age", corsMaxAgeHeader) - if allowOrigin != "*" { - w.Header().Set("Access-Control-Allow-Credentials", "true") - } - vary := w.Header().Get("Vary") - if vary == "" { - w.Header().Set("Vary", "Origin") - } else if !strings.Contains(vary, "Origin") { - w.Header().Set("Vary", vary+", Origin") - } - } else { - log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端=%s", - origin, r.URL.Path, r.Method, GetClientIP(r)) - if isPreflight { - w.WriteHeader(http.StatusForbidden) - return - } - } + // 无 Origin 头:非跨域请求,跳过 CORS 处理 + if origin == "" { + next.ServeHTTP(w, r) + return } + // 有 Origin 头时,响应必须携带 Vary: Origin 防止 CDN 缓存污染 + vary := w.Header().Get("Vary") + if vary == "" { + w.Header().Set("Vary", "Origin") + } else if !strings.Contains(vary, "Origin") { + w.Header().Set("Vary", vary+", Origin") + } + + isPreflight := r.Method == http.MethodOptions + + allowOrigin, matched := matchOrigin(origin) + if !matched { + // Origin 不在白名单:拒绝请求(预检和非预检均拒绝), + // 防止不匹配的请求穿透到后端浪费 TTS 资源 + log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端=%s", + origin, r.URL.Path, r.Method, GetClientIP(r)) + w.WriteHeader(http.StatusForbidden) + return + } + + // Origin 匹配:设置 CORS 响应头 + w.Header().Set("Access-Control-Allow-Origin", allowOrigin) + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") + w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization") + w.Header().Set("Access-Control-Expose-Headers", "X-Request-Id") + w.Header().Set("Access-Control-Max-Age", corsMaxAgeHeader) + if allowOrigin != "*" { + w.Header().Set("Access-Control-Allow-Credentials", "true") + } + + // 预检请求:直接返回 204,不进入内层中间件链, + // 避免消耗速率限制配额和并发槽位 if isPreflight { w.WriteHeader(http.StatusNoContent) return diff --git a/middleware/logger.go b/middleware/logger.go index 6b80991..ea53db2 100644 --- a/middleware/logger.go +++ b/middleware/logger.go @@ -18,13 +18,6 @@ func (rec *statusRecorder) WriteHeader(code int) { func Logger(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/health" { - start := time.Now() - next.ServeHTTP(w, r) - log.Printf("%s %s %s %v", r.Method, r.RequestURI, r.RemoteAddr, time.Since(start)) - return - } - start := time.Now() rec := &statusRecorder{ResponseWriter: w, statusCode: http.StatusOK} next.ServeHTTP(rec, r) diff --git a/middleware/ratelimit.go b/middleware/ratelimit.go index 8d06022..a3ba367 100644 --- a/middleware/ratelimit.go +++ b/middleware/ratelimit.go @@ -90,26 +90,63 @@ func (rl *RateLimiter) cleanup() { } } +// 私有网络 CIDR 范围:仅在直连来源属于这些范围时才信任代理头 +var privateCIDRs []*net.IPNet + +func init() { + for _, cidr := range []string{ + "10.0.0.0/8", + "172.16.0.0/12", + "192.168.0.0/16", + "127.0.0.0/8", + "169.254.0.0/16", + "::1/128", + "fc00::/7", + "fe80::/10", + } { + _, ipNet, _ := net.ParseCIDR(cidr) + privateCIDRs = append(privateCIDRs, ipNet) + } +} + +func isPrivateIP(ipStr string) bool { + ip := net.ParseIP(ipStr) + if ip == nil { + return false + } + if ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() { + return true + } + for _, cidr := range privateCIDRs { + if cidr.Contains(ip) { + return true + } + } + return false +} + +// GetClientIP 提取客户端真实 IP。 +// 仅当直连来源为私有网络(本地代理、Docker 网桥等)时才信任 X-Forwarded-For / X-Real-IP, +// 防止公网直连场景下攻击者伪造代理头绕过速率限制。 func GetClientIP(r *http.Request) string { - xForwardedFor := r.Header.Get("X-Forwarded-For") - if xForwardedFor != "" { - ips := strings.Split(xForwardedFor, ",") - if len(ips) > 0 { - ip := strings.TrimSpace(ips[0]) - if ip != "" { + directIP, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + directIP = r.RemoteAddr + } + + if isPrivateIP(directIP) { + if xff := r.Header.Get("X-Forwarded-For"); xff != "" { + ip := strings.TrimSpace(strings.Split(xff, ",")[0]) + if net.ParseIP(ip) != nil { return ip } } + if xri := strings.TrimSpace(r.Header.Get("X-Real-IP")); xri != "" { + if net.ParseIP(xri) != nil { + return xri + } + } } - xRealIP := strings.TrimSpace(r.Header.Get("X-Real-IP")) - if xRealIP != "" { - return xRealIP - } - - host, _, err := net.SplitHostPort(r.RemoteAddr) - if err != nil { - return r.RemoteAddr - } - return host + return directIP } diff --git a/router/router.go b/router/router.go index d7a4e4c..4635c0f 100644 --- a/router/router.go +++ b/router/router.go @@ -11,7 +11,6 @@ import ( func Setup() *mux.Router { r := mux.NewRouter() - r.Use(middleware.CORS) r.Use(middleware.SecurityHeaders) r.Use(middleware.RateLimit) r.Use(middleware.ConcurrencyLimit) diff --git a/service/stats.go b/service/stats.go index 3ac75a2..1fce766 100644 --- a/service/stats.go +++ b/service/stats.go @@ -1,8 +1,8 @@ package service import ( - "fmt" "runtime" + "strings" "sync" "time" @@ -10,15 +10,17 @@ import ( ) type Stats struct { - totalRequests int64 - successfulRequests int64 - failedRequests int64 - totalResponseTime time.Duration - recentResponseTimes []float64 - responseTimesIndex int - lastErrors []string - errorsIndex int - mutex sync.RWMutex + totalRequests int64 + successfulRequests int64 + failedRequests int64 + totalResponseTime time.Duration + recentResponseTimes []float64 + responseTimesIndex int + responseTimesCount int + lastErrors []string + errorsIndex int + errorsCount int + mutex sync.RWMutex } var GlobalStats *Stats @@ -39,15 +41,34 @@ func (s *Stats) AddRequest(success bool, responseTime time.Duration, errMsg stri s.recentResponseTimes[s.responseTimesIndex] = responseTime.Seconds() * 1000 s.responseTimesIndex = (s.responseTimesIndex + 1) % common.MaxResponseTimes + if s.responseTimesCount < common.MaxResponseTimes { + s.responseTimesCount++ + } if success { s.successfulRequests++ } else { s.failedRequests++ if errMsg != "" { - errInfo := fmt.Sprintf("%s: %s", time.Now().Format(time.RFC3339), errMsg) - s.lastErrors[s.errorsIndex] = errInfo + now := time.Now().Format(time.RFC3339) + + // 去重:如果最近一条错误的消息内容相同,仅更新时间戳 + if s.errorsCount > 0 { + lastIdx := (s.errorsIndex - 1 + common.MaxErrors) % common.MaxErrors + lastEntry := s.lastErrors[lastIdx] + if sepIdx := strings.Index(lastEntry, ": "); sepIdx != -1 { + if lastEntry[sepIdx+2:] == errMsg { + s.lastErrors[lastIdx] = now + ": " + errMsg + return + } + } + } + + s.lastErrors[s.errorsIndex] = now + ": " + errMsg s.errorsIndex = (s.errorsIndex + 1) % common.MaxErrors + if s.errorsCount < common.MaxErrors { + s.errorsCount++ + } } } } @@ -62,19 +83,32 @@ func (s *Stats) GetSnapshot() (totalRequests int64, successfulRequests int64, fa failedRequests = s.failedRequests totalResponseTime = s.totalResponseTime - recentResponseTimes = make([]float64, 0, common.MaxResponseTimes) - for _, t := range s.recentResponseTimes { - if t > 0 { - recentResponseTimes = append(recentResponseTimes, t) + // 按时间顺序(从旧到新)遍历响应时间环形缓冲区 + recentResponseTimes = make([]float64, 0, s.responseTimesCount) + if s.responseTimesCount > 0 { + start := 0 + if s.responseTimesCount == common.MaxResponseTimes { + start = s.responseTimesIndex + } + for i := 0; i < s.responseTimesCount; i++ { + idx := (start + i) % common.MaxResponseTimes + recentResponseTimes = append(recentResponseTimes, s.recentResponseTimes[idx]) } } - lastErrors = make([]string, 0, common.MaxErrors) - for _, e := range s.lastErrors { - if e != "" { - lastErrors = append(lastErrors, e) + // 按时间顺序(从旧到新)遍历错误环形缓冲区 + lastErrors = make([]string, 0, s.errorsCount) + if s.errorsCount > 0 { + start := 0 + if s.errorsCount == common.MaxErrors { + start = s.errorsIndex + } + for i := 0; i < s.errorsCount; i++ { + idx := (start + i) % common.MaxErrors + lastErrors = append(lastErrors, s.lastErrors[idx]) } } + return }