refactor: 重构并优化项目多项功能
1. 调整CORS中间件挂载位置,重构CORS处理逻辑 2. 重写IP获取逻辑,增加私有网络IP信任校验 3. 优化日志中间件,移除/health接口单独日志逻辑 4. 改进API密钥未配置时的提示信息 5. 重构volcano TTS调用,新增voice参数支持 6. 优化请求体过大错误处理 7. 完善统计服务,修复环形缓冲区遍历逻辑,新增去重错误日志功能
This commit is contained in:
@@ -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,
|
||||
|
||||
+2
-1
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
+2
-1
@@ -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 环境变量(多个密钥用逗号分隔)")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+36
-26
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
+53
-16
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
+54
-20
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user