Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3d50b6c69d | ||
|
|
977e9ccadb |
@@ -39,22 +39,6 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
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)
|
body, err := io.ReadAll(r.Body)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
module github.com/volcano-tts/tts-api
|
module github.com/volcano-tts/tts-api
|
||||||
|
|
||||||
go 1.19
|
go 1.26
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
@@ -30,10 +30,10 @@ func main() {
|
|||||||
|
|
||||||
setting.TTSConfigErr = setting.InitTTSConfig()
|
setting.TTSConfigErr = setting.InitTTSConfig()
|
||||||
if setting.TTSConfigErr != nil {
|
if setting.TTSConfigErr != nil {
|
||||||
log.Printf("警告: 配置初始化失败: %v", setting.TTSConfigErr)
|
log.Printf("警告:配置初始化失败: %v", setting.TTSConfigErr)
|
||||||
log.Printf("服务将继续运行,但TTS功能不可用,请检查环境变量配置")
|
log.Printf("服务将继续运行,但TTS功能不可用,请检查环境变量配置
|
||||||
} else {
|
} else {
|
||||||
log.Printf("配置初始化成功")
|
log.Printf("閰嶇疆鍒濆鍖栨垚鍔?)
|
||||||
}
|
}
|
||||||
|
|
||||||
controller.SetStartTime(time.Now())
|
controller.SetStartTime(time.Now())
|
||||||
|
|||||||
+3
-3
@@ -1,4 +1,4 @@
|
|||||||
package middleware
|
package middleware
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/subtle"
|
"crypto/subtle"
|
||||||
@@ -18,9 +18,9 @@ func InitAPIKeys() {
|
|||||||
for i, k := range validAPIKeys {
|
for i, k := range validAPIKeys {
|
||||||
validAPIKeys[i] = strings.TrimSpace(k)
|
validAPIKeys[i] = strings.TrimSpace(k)
|
||||||
}
|
}
|
||||||
log.Printf("已配置 %d 个有效的API密钥", len(validAPIKeys))
|
log.Printf("宸查厤缃?%d 涓湁鏁堢殑API瀵嗛挜", len(validAPIKeys))
|
||||||
} else {
|
} else {
|
||||||
log.Println("警告: OPENAI_TTS_API_KEY 环境变量未设置,将允许所有请求")
|
log.Println("警告: OPENAI_TTS_API_KEY环境变量未设置,将拒绝所有请求?)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+10
-14
@@ -1,4 +1,4 @@
|
|||||||
package middleware
|
package middleware
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"log"
|
"log"
|
||||||
@@ -8,8 +8,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
allowedOrigins []string
|
allowedOrigins []string
|
||||||
allowAllOrigins bool
|
allowAllOrigins bool
|
||||||
corsMaxAgeHeader = "86400"
|
corsMaxAgeHeader = "86400"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -46,9 +46,6 @@ func InitCORSConfig() {
|
|||||||
}
|
}
|
||||||
if len(allowedOrigins) > 0 {
|
if len(allowedOrigins) > 0 {
|
||||||
log.Printf("已配置 %d 个允许的跨域来源白名单", len(allowedOrigins))
|
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" {
|
if origin == "" || origin == "null" || origin == "nil" {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
// 使用小写比较,避免大小写问题
|
|
||||||
lowerOrigin := strings.ToLower(origin)
|
lowerOrigin := strings.ToLower(origin)
|
||||||
if !strings.HasPrefix(lowerOrigin, "http://") && !strings.HasPrefix(lowerOrigin, "https://") {
|
if !strings.HasPrefix(lowerOrigin, "http://") && !strings.HasPrefix(lowerOrigin, "https://") {
|
||||||
return false
|
return false
|
||||||
@@ -66,28 +62,24 @@ func isValidOrigin(origin string) bool {
|
|||||||
|
|
||||||
func matchOrigin(origin string) (string, bool) {
|
func matchOrigin(origin string) (string, bool) {
|
||||||
if !isValidOrigin(origin) {
|
if !isValidOrigin(origin) {
|
||||||
log.Printf("[CORS] Origin %q 验证失败", origin)
|
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
if allowAllOrigins {
|
if allowAllOrigins {
|
||||||
log.Printf("[CORS] Origin %q 匹配 allowAllOrigins", origin)
|
|
||||||
return "*", true
|
return "*", true
|
||||||
}
|
}
|
||||||
normalized := normalizeOrigin(origin)
|
normalized := normalizeOrigin(origin)
|
||||||
log.Printf("[CORS] 检查 origin %q (normalized: %q) 对比白名单: %v", origin, normalized, allowedOrigins)
|
|
||||||
for _, allowed := range allowedOrigins {
|
for _, allowed := range allowedOrigins {
|
||||||
if allowed == normalized {
|
if allowed == normalized {
|
||||||
log.Printf("[CORS] Origin %q 匹配白名单 %q", origin, allowed)
|
|
||||||
return origin, true
|
return origin, true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
log.Printf("[CORS] Origin %q 未匹配任何白名单", origin)
|
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
func CORS(next http.Handler) http.Handler {
|
func CORS(next http.Handler) http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
origin := r.Header.Get("Origin")
|
origin := r.Header.Get("Origin")
|
||||||
|
isPreflight := r.Method == http.MethodOptions
|
||||||
|
|
||||||
if origin != "" {
|
if origin != "" {
|
||||||
allowOrigin, matched := matchOrigin(origin)
|
allowOrigin, matched := matchOrigin(origin)
|
||||||
@@ -107,12 +99,16 @@ func CORS(next http.Handler) http.Handler {
|
|||||||
w.Header().Set("Vary", vary+", Origin")
|
w.Header().Set("Vary", vary+", Origin")
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端IP=%s",
|
log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端=%s",
|
||||||
origin, r.URL.Path, r.Method, GetClientIP(r))
|
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)
|
w.WriteHeader(http.StatusNoContent)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -13,6 +13,8 @@ func Setup() *mux.Router {
|
|||||||
|
|
||||||
r.Use(middleware.CORS)
|
r.Use(middleware.CORS)
|
||||||
r.Use(middleware.SecurityHeaders)
|
r.Use(middleware.SecurityHeaders)
|
||||||
|
r.Use(middleware.RateLimit)
|
||||||
|
r.Use(middleware.ConcurrencyLimit)
|
||||||
r.Use(middleware.Logger)
|
r.Use(middleware.Logger)
|
||||||
|
|
||||||
r.HandleFunc("/v1/audio/speech", controller.OpenaiTTSHandler).Methods("POST", "OPTIONS")
|
r.HandleFunc("/v1/audio/speech", controller.OpenaiTTSHandler).Methods("POST", "OPTIONS")
|
||||||
|
|||||||
Reference in New Issue
Block a user