From 4c936382507045d3714a36ded5a2d917dc2df33c Mon Sep 17 00:00:00 2001 From: sun <3371392206@qq.com> Date: Wed, 20 May 2026 22:26:21 +0800 Subject: [PATCH] =?UTF-8?q?feat(tts=20server):=20=E6=96=B0=E5=A2=9E?= =?UTF-8?q?=E8=B7=A8=E5=9F=9F=E7=99=BD=E5=90=8D=E5=8D=95=E9=85=8D=E7=BD=AE?= =?UTF-8?q?=E5=B9=B6=E4=BC=98=E5=8C=96=E9=94=99=E8=AF=AF=E5=A4=84=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1. 重构CORS中间件,支持通过ALLOWED_ORIGINS环境变量配置跨域白名单,默认允许所有来源并添加凭证支持 2. 优化TTS服务错误响应体读取逻辑,处理读取失败的情况 3. 精简健康检查接口返回的冗余配置信息 4. 删除部分重复的配置日志输出 --- tts_server.go | 56 ++++++++++++++++++++++++++++++++++++++++----------- 1 file changed, 44 insertions(+), 12 deletions(-) diff --git a/tts_server.go b/tts_server.go index d053444..4d4159d 100644 --- a/tts_server.go +++ b/tts_server.go @@ -333,8 +333,12 @@ func synthesis(text string, speed float64) (*SynthesisResult, error) { defer resp.Body.Close() if resp.StatusCode != http.StatusOK { - body, _ := io.ReadAll(resp.Body) - log.Printf("TTS service error: status=%d, body=%s", resp.StatusCode, string(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)) + } 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{}{ "recent_errors_count": len(lastErrors), - "recent_errors": lastErrors, }, "config_status": map[string]interface{}{ "all_required_vars_set": allEnvVarsSet, "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) } +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 { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Access-Control-Allow-Origin", "*") - w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") - w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization") + 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-Headers", "Content-Type, Authorization") + w.Header().Set("Access-Control-Allow-Credentials", "true") + } if r.Method == http.MethodOptions { w.WriteHeader(http.StatusOK) @@ -700,6 +733,7 @@ func main() { log.SetPrefix("[TTS-Server] ") initAPIKeys() + initCORSConfig() ttsConfigErr = initTTSConfig() if ttsConfigErr != nil { @@ -761,9 +795,7 @@ func main() { log.Printf("Listening on port: %s", port) log.Printf("OpenAI TTS endpoint: http://localhost:%s/v1/audio/speech", port) log.Printf("Health check: http://localhost:%s/health", port) - log.Printf("Using ByteDance v3 API: %s", ttsConfig.URL) - log.Printf("Resource ID: %s", ttsConfig.ResourceId) - log.Printf("Speaker: %s", ttsConfig.Speaker) + log.Printf("Using ByteDance v3 API") if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed { log.Fatalf("Server failed to start: %v", err)