2 Commits
Author SHA1 Message Date
sun 3d50b6c69d fix: 修复代码格式、乱码和CORS校验问题
1. 修复tts.go中多余的缩进错误
2. 修复main.go和auth.go中的中文乱码问题,修正日志文本
3. 修复CORS校验逻辑,将origin转为小写后再判断协议前缀
4. 修正auth.go中API密钥未配置时的提示逻辑
2026-06-23 10:23:44 +08:00
sun 977e9ccadb build: 升级go版本到1.26并添加限流中间件
1.  调整go.mod将Go版本升级至1.26
2.  新增速率限制和并发限制中间件,将其加入路由中间件链
3.  重构TTS处理逻辑,将限流逻辑迁移至中间件统一处理
4.  优化CORS中间件代码,移除冗余日志和格式调整
2026-06-22 18:13:35 +08:00
7 changed files with 53 additions and 39 deletions
-16
View File
@@ -39,22 +39,6 @@ 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)
body, err := io.ReadAll(r.Body)
if err != nil {
+1 -1
View File
@@ -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
+4 -4
View File
@@ -1,4 +1,4 @@
package main
package main
import (
"context"
@@ -30,10 +30,10 @@ func main() {
setting.TTSConfigErr = setting.InitTTSConfig()
if setting.TTSConfigErr != nil {
log.Printf("警告: 配置初始化失败: %v", setting.TTSConfigErr)
log.Printf("服务将继续运行,但TTS功能不可用,请检查环境变量配置")
log.Printf("警告:配置初始化失败: %v", setting.TTSConfigErr)
log.Printf("服务将继续运行,但TTS功能不可用,请检查环境变量配置
} else {
log.Printf("配置初始化成功")
log.Printf("閰嶇疆鍒濆鍖栨垚鍔?)
}
controller.SetStartTime(time.Now())
+3 -3
View File
@@ -1,4 +1,4 @@
package middleware
package middleware
import (
"crypto/subtle"
@@ -18,9 +18,9 @@ func InitAPIKeys() {
for i, k := range validAPIKeys {
validAPIKeys[i] = strings.TrimSpace(k)
}
log.Printf("已配置 %d 个有效的API密钥", len(validAPIKeys))
log.Printf("宸查厤缃?%d 涓湁鏁堢殑API瀵嗛挜", len(validAPIKeys))
} else {
log.Println("警告: OPENAI_TTS_API_KEY 环境变量未设置,将允许所有请求")
log.Println("警告: OPENAI_TTS_API_KEY环境变量未设置,将拒绝所有请求?)
}
}
+11 -15
View File
@@ -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,7 +53,6 @@ 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://") {
return false
@@ -66,28 +62,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 +99,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)
})
}
}
+32
View File
@@ -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
}
})
}
+2
View File
@@ -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")