From 977e9ccadbc1d38a1ef633b2002bebe19a5ba453 Mon Sep 17 00:00:00 2001 From: sun <3371392206@qq.com> Date: Mon, 22 Jun 2026 18:13:35 +0800 Subject: [PATCH] =?UTF-8?q?build:=20=E5=8D=87=E7=BA=A7go=E7=89=88=E6=9C=AC?= =?UTF-8?q?=E5=88=B01.26=E5=B9=B6=E6=B7=BB=E5=8A=A0=E9=99=90=E6=B5=81?= =?UTF-8?q?=E4=B8=AD=E9=97=B4=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1. 调整go.mod将Go版本升级至1.26 2. 新增速率限制和并发限制中间件,将其加入路由中间件链 3. 重构TTS处理逻辑,将限流逻辑迁移至中间件统一处理 4. 优化CORS中间件代码,移除冗余日志和格式调整 --- controller/tts.go | 18 +---------------- go.mod | 2 +- middleware/cors.go | 29 +++++++++++---------------- middleware/ratelimit_middleware.go | 32 ++++++++++++++++++++++++++++++ router/router.go | 2 ++ 5 files changed, 48 insertions(+), 35 deletions(-) create mode 100644 middleware/ratelimit_middleware.go diff --git a/controller/tts.go b/controller/tts.go index 03c6297..8a34057 100644 --- a/controller/tts.go +++ b/controller/tts.go @@ -39,23 +39,7 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) { return } - select { - case middleware.ConcurrencySem <- struct{}{}: - defer func() { <-middleware.ConcurrencySem }() - default: - log.Printf("警告: 已达到最大并发请求数限制,拒绝请求 - 客户端IP: %s", middleware.GetClientIP(r)) - middleware.SendJSONError(w, http.StatusServiceUnavailable, "Server is busy, maximum concurrent requests reached. Please try again later.", "concurrency_limit_error", "max_concurrent_requests") - return - } - - clientIP := middleware.GetClientIP(r) - if !middleware.GlobalRateLimiter.Allow(clientIP) { - log.Printf("警告: 已超过IP速率限制,拒绝请求 - 客户端IP: %s", clientIP) - middleware.SendJSONError(w, http.StatusTooManyRequests, "Rate limit exceeded. Please try again later.", "rate_limit_error", "rate_limit_exceeded") - return - } - - r.Body = http.MaxBytesReader(w, r.Body, common.MaxRequestBodySize) + r.Body = http.MaxBytesReader(w, r.Body, common.MaxRequestBodySize) body, err := io.ReadAll(r.Body) if err != nil { if strings.Contains(err.Error(), "request body too large") { diff --git a/go.mod b/go.mod index 101db9f..0da2762 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/volcano-tts/tts-api -go 1.19 +go 1.26 require ( github.com/google/uuid v1.6.0 diff --git a/middleware/cors.go b/middleware/cors.go index c391d13..5ba1a08 100644 --- a/middleware/cors.go +++ b/middleware/cors.go @@ -1,4 +1,4 @@ -package middleware +package middleware import ( "log" @@ -8,8 +8,8 @@ import ( ) var ( - allowedOrigins []string - allowAllOrigins bool + allowedOrigins []string + allowAllOrigins bool corsMaxAgeHeader = "86400" ) @@ -46,9 +46,6 @@ func InitCORSConfig() { } if len(allowedOrigins) > 0 { log.Printf("已配置 %d 个允许的跨域来源白名单", len(allowedOrigins)) - for _, o := range allowedOrigins { - log.Printf(" - %s", o) - } } } @@ -56,9 +53,7 @@ func isValidOrigin(origin string) bool { if origin == "" || origin == "null" || origin == "nil" { return false } - // 使用小写比较,避免大小写问题 - lowerOrigin := strings.ToLower(origin) - if !strings.HasPrefix(lowerOrigin, "http://") && !strings.HasPrefix(lowerOrigin, "https://") { + if !strings.HasPrefix(origin, "http://") && !strings.HasPrefix(origin, "https://") { return false } return true @@ -66,28 +61,24 @@ func isValidOrigin(origin string) bool { func matchOrigin(origin string) (string, bool) { if !isValidOrigin(origin) { - log.Printf("[CORS] Origin %q 验证失败", origin) return "", false } if allowAllOrigins { - log.Printf("[CORS] Origin %q 匹配 allowAllOrigins", origin) return "*", true } normalized := normalizeOrigin(origin) - log.Printf("[CORS] 检查 origin %q (normalized: %q) 对比白名单: %v", origin, normalized, allowedOrigins) for _, allowed := range allowedOrigins { if allowed == normalized { - log.Printf("[CORS] Origin %q 匹配白名单 %q", origin, allowed) return origin, true } } - log.Printf("[CORS] Origin %q 未匹配任何白名单", origin) return "", false } 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) @@ -107,16 +98,20 @@ func CORS(next http.Handler) http.Handler { w.Header().Set("Vary", vary+", Origin") } } else { - log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端IP=%s", + log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端=%s", origin, r.URL.Path, r.Method, GetClientIP(r)) + if isPreflight { + w.WriteHeader(http.StatusForbidden) + return + } } } - if r.Method == http.MethodOptions { + if isPreflight { w.WriteHeader(http.StatusNoContent) return } next.ServeHTTP(w, r) }) -} \ No newline at end of file +} diff --git a/middleware/ratelimit_middleware.go b/middleware/ratelimit_middleware.go new file mode 100644 index 0000000..bbbb9f0 --- /dev/null +++ b/middleware/ratelimit_middleware.go @@ -0,0 +1,32 @@ +package middleware + +import ( + "log" + "net/http" +) + +func RateLimit(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + clientIP := GetClientIP(r) + if !GlobalRateLimiter.Allow(clientIP) { + log.Printf("警告: 已超过IP速率限制,拒绝请求 - 客户端IP: %s", clientIP) + SendJSONError(w, http.StatusTooManyRequests, "Rate limit exceeded. Please try again later.", "rate_limit_error", "rate_limit_exceeded") + return + } + next.ServeHTTP(w, r) + }) +} + +func ConcurrencyLimit(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + select { + case ConcurrencySem <- struct{}{}: + defer func() { <-ConcurrencySem }() + next.ServeHTTP(w, r) + default: + log.Printf("警告: 已达到最大并发请求数限制,拒绝请求 - 客户端IP: %s", GetClientIP(r)) + SendJSONError(w, http.StatusServiceUnavailable, "Server is busy, maximum concurrent requests reached. Please try again later.", "concurrency_limit_error", "max_concurrent_requests") + return + } + }) +} diff --git a/router/router.go b/router/router.go index d1c2010..d7a4e4c 100644 --- a/router/router.go +++ b/router/router.go @@ -13,6 +13,8 @@ func Setup() *mux.Router { r.Use(middleware.CORS) r.Use(middleware.SecurityHeaders) + r.Use(middleware.RateLimit) + r.Use(middleware.ConcurrencyLimit) r.Use(middleware.Logger) r.HandleFunc("/v1/audio/speech", controller.OpenaiTTSHandler).Methods("POST", "OPTIONS")