feat(tts server): 新增跨域白名单配置并优化错误处理
1. 重构CORS中间件,支持通过ALLOWED_ORIGINS环境变量配置跨域白名单,默认允许所有来源并添加凭证支持 2. 优化TTS服务错误响应体读取逻辑,处理读取失败的情况 3. 精简健康检查接口返回的冗余配置信息 4. 删除部分重复的配置日志输出
This commit is contained in:
+41
-9
@@ -333,8 +333,12 @@ func synthesis(text string, speed float64) (*SynthesisResult, error) {
|
|||||||
defer resp.Body.Close()
|
defer resp.Body.Close()
|
||||||
|
|
||||||
if resp.StatusCode != http.StatusOK {
|
if resp.StatusCode != http.StatusOK {
|
||||||
body, _ := io.ReadAll(resp.Body)
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
log.Printf("Failed to read error response body: %v", err)
|
||||||
|
} else {
|
||||||
log.Printf("TTS service error: status=%d, body=%s", resp.StatusCode, string(body))
|
log.Printf("TTS service error: status=%d, body=%s", resp.StatusCode, string(body))
|
||||||
|
}
|
||||||
return nil, fmt.Errorf("TTS service error: status %d", resp.StatusCode)
|
return nil, fmt.Errorf("TTS service error: status %d", resp.StatusCode)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -652,14 +656,10 @@ func healthHandler(w http.ResponseWriter, r *http.Request) {
|
|||||||
},
|
},
|
||||||
"errors": map[string]interface{}{
|
"errors": map[string]interface{}{
|
||||||
"recent_errors_count": len(lastErrors),
|
"recent_errors_count": len(lastErrors),
|
||||||
"recent_errors": lastErrors,
|
|
||||||
},
|
},
|
||||||
"config_status": map[string]interface{}{
|
"config_status": map[string]interface{}{
|
||||||
"all_required_vars_set": allEnvVarsSet,
|
"all_required_vars_set": allEnvVarsSet,
|
||||||
"config_error": ttsConfigErr != nil,
|
"config_error": ttsConfigErr != nil,
|
||||||
"config_error_message": fmt.Sprintf("%v", ttsConfigErr),
|
|
||||||
"resource_id": ttsConfig.ResourceId,
|
|
||||||
"speaker": ttsConfig.Speaker,
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -678,11 +678,44 @@ func (rec *statusRecorder) WriteHeader(code int) {
|
|||||||
rec.ResponseWriter.WriteHeader(code)
|
rec.ResponseWriter.WriteHeader(code)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var allowedOrigins []string
|
||||||
|
|
||||||
|
func initCORSConfig() {
|
||||||
|
origins := os.Getenv("ALLOWED_ORIGINS")
|
||||||
|
if origins != "" {
|
||||||
|
allowedOrigins = strings.Split(origins, ",")
|
||||||
|
for i, origin := range allowedOrigins {
|
||||||
|
allowedOrigins[i] = strings.TrimSpace(origin)
|
||||||
|
}
|
||||||
|
log.Printf("已配置 %d 个允许的跨域来源", len(allowedOrigins))
|
||||||
|
} else {
|
||||||
|
log.Println("警告: ALLOWED_ORIGINS 环境变量未设置")
|
||||||
|
log.Println("将使用反射模式允许所有请求来源,生产环境建议配置白名单")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func corsMiddleware(next http.Handler) http.Handler {
|
func corsMiddleware(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) {
|
||||||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
origin := r.Header.Get("Origin")
|
||||||
|
allowOrigin := ""
|
||||||
|
|
||||||
|
if len(allowedOrigins) == 0 {
|
||||||
|
allowOrigin = origin
|
||||||
|
} else {
|
||||||
|
for _, allowed := range allowedOrigins {
|
||||||
|
if allowed == origin {
|
||||||
|
allowOrigin = origin
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if 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")
|
||||||
|
w.Header().Set("Access-Control-Allow-Credentials", "true")
|
||||||
|
}
|
||||||
|
|
||||||
if r.Method == http.MethodOptions {
|
if r.Method == http.MethodOptions {
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
@@ -700,6 +733,7 @@ func main() {
|
|||||||
log.SetPrefix("[TTS-Server] ")
|
log.SetPrefix("[TTS-Server] ")
|
||||||
|
|
||||||
initAPIKeys()
|
initAPIKeys()
|
||||||
|
initCORSConfig()
|
||||||
|
|
||||||
ttsConfigErr = initTTSConfig()
|
ttsConfigErr = initTTSConfig()
|
||||||
if ttsConfigErr != nil {
|
if ttsConfigErr != nil {
|
||||||
@@ -761,9 +795,7 @@ func main() {
|
|||||||
log.Printf("Listening on port: %s", port)
|
log.Printf("Listening on port: %s", port)
|
||||||
log.Printf("OpenAI TTS endpoint: http://localhost:%s/v1/audio/speech", port)
|
log.Printf("OpenAI TTS endpoint: http://localhost:%s/v1/audio/speech", port)
|
||||||
log.Printf("Health check: http://localhost:%s/health", port)
|
log.Printf("Health check: http://localhost:%s/health", port)
|
||||||
log.Printf("Using ByteDance v3 API: %s", ttsConfig.URL)
|
log.Printf("Using ByteDance v3 API")
|
||||||
log.Printf("Resource ID: %s", ttsConfig.ResourceId)
|
|
||||||
log.Printf("Speaker: %s", ttsConfig.Speaker)
|
|
||||||
|
|
||||||
if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||||||
log.Fatalf("Server failed to start: %v", err)
|
log.Fatalf("Server failed to start: %v", err)
|
||||||
|
|||||||
Reference in New Issue
Block a user