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"
|
"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
@@ -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 {
|
||||||
|
|||||||
@@ -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
@@ -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 环境变量(多个密钥用逗号分隔)")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+27
-17
@@ -79,11 +79,34 @@ 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")
|
||||||
|
|
||||||
|
// 无 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
|
isPreflight := r.Method == http.MethodOptions
|
||||||
|
|
||||||
if origin != "" {
|
|
||||||
allowOrigin, matched := matchOrigin(origin)
|
allowOrigin, matched := matchOrigin(origin)
|
||||||
if matched {
|
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-Origin", allowOrigin)
|
||||||
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
|
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-Allow-Headers", "Content-Type, Authorization")
|
||||||
@@ -92,22 +115,9 @@ func CORS(next http.Handler) http.Handler {
|
|||||||
if allowOrigin != "*" {
|
if allowOrigin != "*" {
|
||||||
w.Header().Set("Access-Control-Allow-Credentials", "true")
|
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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
|
// 预检请求:直接返回 204,不进入内层中间件链,
|
||||||
|
// 避免消耗速率限制配额和并发槽位
|
||||||
if isPreflight {
|
if isPreflight {
|
||||||
w.WriteHeader(http.StatusNoContent)
|
w.WriteHeader(http.StatusNoContent)
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -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
@@ -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
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
+45
-11
@@ -1,8 +1,8 @@
|
|||||||
package service
|
package service
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -16,8 +16,10 @@ type Stats struct {
|
|||||||
totalResponseTime time.Duration
|
totalResponseTime time.Duration
|
||||||
recentResponseTimes []float64
|
recentResponseTimes []float64
|
||||||
responseTimesIndex int
|
responseTimesIndex int
|
||||||
|
responseTimesCount int
|
||||||
lastErrors []string
|
lastErrors []string
|
||||||
errorsIndex int
|
errorsIndex int
|
||||||
|
errorsCount int
|
||||||
mutex sync.RWMutex
|
mutex sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user