refactor: 重构并优化项目多项功能

1. 调整CORS中间件挂载位置,重构CORS处理逻辑
2. 重写IP获取逻辑,增加私有网络IP信任校验
3. 优化日志中间件,移除/health接口单独日志逻辑
4. 改进API密钥未配置时的提示信息
5. 重构volcano TTS调用,新增voice参数支持
6. 优化请求体过大错误处理
7. 完善统计服务,修复环形缓冲区遍历逻辑,新增去重错误日志功能
This commit is contained in:
sun
2026-06-26 19:12:50 +08:00
parent 3d50b6c69d
commit 61431e00ba
9 changed files with 161 additions and 82 deletions
+13 -9
View File
@@ -13,7 +13,6 @@ import (
"time" "time"
"github.com/google/uuid" "github.com/google/uuid"
"github.com/volcano-tts/tts-api/common"
"github.com/volcano-tts/tts-api/dto" "github.com/volcano-tts/tts-api/dto"
) )
@@ -24,7 +23,6 @@ type HTTPClient struct {
func NewHTTPClient() *HTTPClient { func NewHTTPClient() *HTTPClient {
return &HTTPClient{ return &HTTPClient{
client: &http.Client{ client: &http.Client{
Timeout: common.DefaultTimeout,
Transport: &http.Transport{ Transport: &http.Transport{
MaxIdleConns: 100, MaxIdleConns: 100,
MaxIdleConnsPerHost: 20, MaxIdleConnsPerHost: 20,
@@ -52,19 +50,25 @@ func (h *HTTPClient) PostStream(url string, headers map[string]string, body []by
} }
func convertSpeedToSpeechRate(speed float64) int { func convertSpeedToSpeechRate(speed float64) int {
if speed <= 0.5 { rate := int((speed - 1.0) * 100)
return -50 if rate < -200 {
rate = -200
} }
if speed >= 2.0 { if rate > 500 {
return 100 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() reqID := uuid.NewString()
speechRate := convertSpeedToSpeechRate(speed) speechRate := convertSpeedToSpeechRate(speed)
speaker := config.Speaker
if voice != "" {
speaker = voice
}
params := map[string]interface{}{ params := map[string]interface{}{
"user": map[string]interface{}{ "user": map[string]interface{}{
"uid": "uid", "uid": "uid",
@@ -72,7 +76,7 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
"namespace": "BidirectionalTTS", "namespace": "BidirectionalTTS",
"req_params": map[string]interface{}{ "req_params": map[string]interface{}{
"text": text, "text": text,
"speaker": config.Speaker, "speaker": speaker,
"audio_params": map[string]interface{}{ "audio_params": map[string]interface{}{
"format": "wav", "format": "wav",
"sample_rate": 24000, "sample_rate": 24000,
+2 -1
View File
@@ -43,6 +43,7 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
body, err := io.ReadAll(r.Body) body, err := io.ReadAll(r.Body)
if err != nil { if err != nil {
if strings.Contains(err.Error(), "request body too large") { if strings.Contains(err.Error(), "request body too large") {
http.Error(w, "Request body too large", http.StatusRequestEntityTooLarge)
return return
} }
http.Error(w, "Failed to read request body", http.StatusBadRequest) 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() 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) duration := time.Since(ttsStart)
if err != nil { if err != nil {
+1 -1
View File
@@ -47,7 +47,7 @@ func main() {
server := &http.Server{ server := &http.Server{
Addr: ":" + port, Addr: ":" + port,
Handler: r, Handler: middleware.CORS(r),
ReadTimeout: 30 * time.Second, ReadTimeout: 30 * time.Second,
WriteTimeout: 120 * time.Second, WriteTimeout: 120 * time.Second,
IdleTimeout: 60 * time.Second, IdleTimeout: 60 * time.Second,
+2 -1
View File
@@ -20,7 +20,8 @@ func InitAPIKeys() {
} }
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环境变量未设置,所有请求将无需认证即可访问")
log.Println("如需启用API密钥验证,请设置 OPENAI_TTS_API_KEY 环境变量(多个密钥用逗号分隔)")
} }
} }
+36 -26
View File
@@ -79,35 +79,45 @@ func matchOrigin(origin string) (string, bool) {
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 != "" { // 无 Origin 头:非跨域请求,跳过 CORS 处理
allowOrigin, matched := matchOrigin(origin) if origin == "" {
if matched { next.ServeHTTP(w, r)
w.Header().Set("Access-Control-Allow-Origin", allowOrigin) return
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 头时,响应必须携带 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 { if isPreflight {
w.WriteHeader(http.StatusNoContent) w.WriteHeader(http.StatusNoContent)
return return
-7
View File
@@ -18,13 +18,6 @@ func (rec *statusRecorder) WriteHeader(code int) {
func Logger(next http.Handler) http.Handler { func Logger(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) {
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() start := time.Now()
rec := &statusRecorder{ResponseWriter: w, statusCode: http.StatusOK} rec := &statusRecorder{ResponseWriter: w, statusCode: http.StatusOK}
next.ServeHTTP(rec, r) next.ServeHTTP(rec, r)
+53 -16
View File
@@ -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 { func GetClientIP(r *http.Request) string {
xForwardedFor := r.Header.Get("X-Forwarded-For") directIP, _, err := net.SplitHostPort(r.RemoteAddr)
if xForwardedFor != "" { if err != nil {
ips := strings.Split(xForwardedFor, ",") directIP = r.RemoteAddr
if len(ips) > 0 { }
ip := strings.TrimSpace(ips[0])
if ip != "" { 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 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")) return directIP
if xRealIP != "" {
return xRealIP
}
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return r.RemoteAddr
}
return host
} }
-1
View File
@@ -11,7 +11,6 @@ import (
func Setup() *mux.Router { func Setup() *mux.Router {
r := mux.NewRouter() r := mux.NewRouter()
r.Use(middleware.CORS)
r.Use(middleware.SecurityHeaders) r.Use(middleware.SecurityHeaders)
r.Use(middleware.RateLimit) r.Use(middleware.RateLimit)
r.Use(middleware.ConcurrencyLimit) r.Use(middleware.ConcurrencyLimit)
+54 -20
View File
@@ -1,8 +1,8 @@
package service package service
import ( import (
"fmt"
"runtime" "runtime"
"strings"
"sync" "sync"
"time" "time"
@@ -10,15 +10,17 @@ import (
) )
type Stats struct { type Stats struct {
totalRequests int64 totalRequests int64
successfulRequests int64 successfulRequests int64
failedRequests int64 failedRequests int64
totalResponseTime time.Duration totalResponseTime time.Duration
recentResponseTimes []float64 recentResponseTimes []float64
responseTimesIndex int responseTimesIndex int
lastErrors []string responseTimesCount int
errorsIndex int lastErrors []string
mutex sync.RWMutex errorsIndex int
errorsCount int
mutex sync.RWMutex
} }
var GlobalStats *Stats 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.recentResponseTimes[s.responseTimesIndex] = responseTime.Seconds() * 1000
s.responseTimesIndex = (s.responseTimesIndex + 1) % common.MaxResponseTimes s.responseTimesIndex = (s.responseTimesIndex + 1) % common.MaxResponseTimes
if s.responseTimesCount < common.MaxResponseTimes {
s.responseTimesCount++
}
if success { if success {
s.successfulRequests++ s.successfulRequests++
} else { } else {
s.failedRequests++ s.failedRequests++
if errMsg != "" { if errMsg != "" {
errInfo := fmt.Sprintf("%s: %s", time.Now().Format(time.RFC3339), errMsg) now := time.Now().Format(time.RFC3339)
s.lastErrors[s.errorsIndex] = errInfo
// 去重:如果最近一条错误的消息内容相同,仅更新时间戳
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 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 failedRequests = s.failedRequests
totalResponseTime = s.totalResponseTime totalResponseTime = s.totalResponseTime
recentResponseTimes = make([]float64, 0, common.MaxResponseTimes) // 按时间顺序(从旧到新)遍历响应时间环形缓冲区
for _, t := range s.recentResponseTimes { recentResponseTimes = make([]float64, 0, s.responseTimesCount)
if t > 0 { if s.responseTimesCount > 0 {
recentResponseTimes = append(recentResponseTimes, t) 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 { lastErrors = make([]string, 0, s.errorsCount)
if e != "" { if s.errorsCount > 0 {
lastErrors = append(lastErrors, e) 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 return
} }