diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..45b0f8d --- /dev/null +++ b/.dockerignore @@ -0,0 +1,9 @@ +*.exe +*.md +.env +.env.example +.git +.gitignore +tts_api_architecture.html +代码审查报告.md +fix_list.md diff --git a/.env.example b/.env.example index a1e9875..99500f3 100644 --- a/.env.example +++ b/.env.example @@ -9,19 +9,9 @@ BYTEDANCE_TTS_API_KEY=your_api_key_here # 资源信息ID(决定使用1.0还是2.0模型) -# 语音合成模型: -# - seed-tts-1.0: 豆包语音合成模型1.0字符版 -# - seed-tts-1.0-concurr: 豆包语音合成模型1.0并发版 -# - seed-tts-2.0: 豆包语音合成模型2.0字符版 -# 声音复刻模型: -# - seed-icl-1.0: 声音复刻1.0字符版 -# - seed-icl-1.0-concurr: 声音复刻1.0并发版 -# - seed-icl-2.0: 声音复刻2.0字符版 BYTEDANCE_TTS_RESOURCE_ID=seed-tts-1.0 -# 发音人(音色)ID,具体参考火山引擎音色列表 -# 注意:1.0音色只能搭配 seed-tts-1.0 Resource ID -# 2.0音色只能搭配 seed-tts-2.0 Resource ID +# 发音人(音色)ID BYTEDANCE_TTS_SPEAKER=your_speaker_id_here # ========================================== @@ -32,8 +22,10 @@ BYTEDANCE_TTS_SPEAKER=your_speaker_id_here BYTEDANCE_TTS_TIMEOUT=30s # OpenAI兼容接口的API密钥(可选) -# 配置后,客户端请求需要携带 Authorization: Bearer OPENAI_TTS_API_KEY=your_openai_compatible_key_here +# CORS 跨域白名单(逗号分隔,开发环境可设 *) +# ALLOWED_ORIGINS=https://example.com,https://app.example.com + # 服务监听端口,默认8080 PORT=8080 diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..4efb3ba --- /dev/null +++ b/Dockerfile @@ -0,0 +1,26 @@ +FROM golang:1.19-alpine AS builder + +WORKDIR /app + +COPY go.mod go.sum ./ +RUN go mod download + +COPY . . + +RUN CGO_ENABLED=0 GOOS=linux go build -o tts-api . + +FROM alpine:3.18 + +RUN apk --no-cache add ca-certificates tzdata + +WORKDIR /app + +COPY --from=builder /app/tts-api . +COPY --from=builder /app/health.html . + +EXPOSE 8080 + +HEALTHCHECK --interval=30s --timeout=5s --start-period=5s --retries=3 \ + CMD wget -qO- http://localhost:8080/health || exit 1 + +ENTRYPOINT ["./tts-api"] diff --git a/adapter/volcano/volcano.go b/adapter/volcano/volcano.go new file mode 100644 index 0000000..603ff7c --- /dev/null +++ b/adapter/volcano/volcano.go @@ -0,0 +1,167 @@ +package volcano + +import ( + "bufio" + "bytes" + "context" + "encoding/base64" + "encoding/json" + "fmt" + "io" + "log" + "net/http" + "time" + + "github.com/google/uuid" + "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/dto" +) + +type HTTPClient struct { + client *http.Client +} + +func NewHTTPClient() *HTTPClient { + return &HTTPClient{ + client: &http.Client{ + Timeout: common.DefaultTimeout, + Transport: &http.Transport{ + MaxIdleConns: 100, + MaxIdleConnsPerHost: 20, + IdleConnTimeout: 90 * time.Second, + TLSHandshakeTimeout: 10 * time.Second, + }, + }, + } +} + +func (h *HTTPClient) PostStream(url string, headers map[string]string, body []byte, timeout time.Duration) (*http.Response, error) { + req, err := http.NewRequest(http.MethodPost, url, bytes.NewBuffer(body)) + if err != nil { + return nil, err + } + for key, value := range headers { + req.Header.Set(key, value) + } + + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + req = req.WithContext(ctx) + return h.client.Do(req) +} + +func convertSpeedToSpeechRate(speed float64) int { + if speed <= 0.5 { + return -50 + } + if speed >= 2.0 { + return 100 + } + return int((speed - 1.0) * 100) +} + +func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text string, speed float64) (*dto.SynthesisResult, error) { + reqID := uuid.NewString() + speechRate := convertSpeedToSpeechRate(speed) + + params := map[string]interface{}{ + "user": map[string]interface{}{ + "uid": "uid", + }, + "namespace": "BidirectionalTTS", + "req_params": map[string]interface{}{ + "text": text, + "speaker": config.Speaker, + "audio_params": map[string]interface{}{ + "format": "wav", + "sample_rate": 24000, + "speech_rate": speechRate, + }, + }, + } + + headers := map[string]string{ + "Content-Type": "application/json", + "Connection": "keep-alive", + "X-Api-Resource-Id": config.ResourceId, + "X-Api-Request-Id": reqID, + "X-Api-Key": config.ApiKey, + } + + bodyStr, err := json.Marshal(params) + if err != nil { + log.Printf("JSON marshal fail: %v", err) + return nil, err + } + + resp, err := httpClient.PostStream(config.URL, headers, bodyStr, config.Timeout) + if err != nil { + log.Printf("http post fail: %v", err) + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + 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) + } + + var audioData []byte + scanner := bufio.NewScanner(resp.Body) + scanner.Buffer(make([]byte, 1024*1024), 8*1024*1024) + + for scanner.Scan() { + line := scanner.Bytes() + if len(line) == 0 { + continue + } + + var v3Resp dto.V3TTSResponse + if err := json.Unmarshal(line, &v3Resp); err != nil { + log.Printf("unmarshal chunk fail: %v, line: %s", err, string(line)) + continue + } + + if v3Resp.Code == 20000000 { + if v3Resp.Usage != nil { + log.Printf("TTS synthesis completed, usage: %+v", v3Resp.Usage) + } + for scanner.Scan() { + } + break + } + + if v3Resp.Code != 0 { + log.Printf("TTS service error: code=%d, message=%s", v3Resp.Code, v3Resp.Message) + return nil, fmt.Errorf("TTS service error: %s", v3Resp.Message) + } + + if v3Resp.Data != "" { + chunk, err := base64.StdEncoding.DecodeString(v3Resp.Data) + if err != nil { + log.Printf("base64 decode fail: %v", err) + return nil, err + } + audioData = append(audioData, chunk...) + } else if v3Resp.Sentence != "" { + log.Printf("Received sentence info (sequence %d): %s", v3Resp.Sequence, v3Resp.Sentence) + } + } + + if err := scanner.Err(); err != nil { + log.Printf("read stream fail: %v", err) + return nil, err + } + + if len(audioData) == 0 { + return nil, fmt.Errorf("no audio data received") + } + + return &dto.SynthesisResult{AudioData: audioData, ReqID: reqID}, nil +} diff --git a/common/constants.go b/common/constants.go new file mode 100644 index 0000000..493eef7 --- /dev/null +++ b/common/constants.go @@ -0,0 +1,19 @@ +package common + +import "time" + +const ( + DefaultPort = "8080" + DefaultTimeout = 30 * time.Second + MaxTextLength = 5000 + MinSpeed = 0.25 + MaxSpeed = 4.0 + DefaultSpeed = 1.0 + MaxRequestBodySize = 1024 * 1024 + RateLimitRequests = 100 + RateLimitWindow = time.Minute + MaxResponseTimes = 100 + MaxErrors = 10 + MaxConcurrentRequests = 10 + CleanupInterval = time.Hour +) diff --git a/controller/tts.go b/controller/tts.go new file mode 100644 index 0000000..9185442 --- /dev/null +++ b/controller/tts.go @@ -0,0 +1,174 @@ +package controller + +import ( + "encoding/json" + "fmt" + "io" + "log" + "net/http" + "strings" + "time" + + "github.com/volcano-tts/tts-api/adapter/volcano" + "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/dto" + "github.com/volcano-tts/tts-api/middleware" + "github.com/volcano-tts/tts-api/service" + "github.com/volcano-tts/tts-api/setting" +) + +var volcanoClient *volcano.HTTPClient + +func InitController() { + volcanoClient = volcano.NewHTTPClient() +} + +func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) + return + } + + if !middleware.ValidateAPIKey(r) { + middleware.SendJSONError(w, http.StatusUnauthorized, "Invalid API key provided.", "invalid_request_error", "invalid_api_key") + return + } + + if setting.TTSConfigErr != nil { + middleware.SendJSONError(w, http.StatusServiceUnavailable, "TTS service configuration error. Please check environment variables and restart the service.", "configuration_error", "service_unavailable") + return + } + + select { + case middleware.ConcurrencySem <- struct{}{}: + defer func() { <-middleware.ConcurrencySem }() + default: + log.Printf("警告: 已达到最大并发请求数限制,拒绝请求 - 客户端IP: %s", middleware.GetClientIP(r)) + middleware.SendJSONError(w, http.StatusServiceUnavailable, "Server is busy, maximum concurrent requests reached. Please try again later.", "concurrency_limit_error", "max_concurrent_requests") + return + } + + clientIP := middleware.GetClientIP(r) + if !middleware.GlobalRateLimiter.Allow(clientIP) { + log.Printf("警告: 已超过IP速率限制,拒绝请求 - 客户端IP: %s", clientIP) + middleware.SendJSONError(w, http.StatusTooManyRequests, "Rate limit exceeded. Please try again later.", "rate_limit_error", "rate_limit_exceeded") + return + } + + r.Body = http.MaxBytesReader(w, r.Body, common.MaxRequestBodySize) + body, err := io.ReadAll(r.Body) + if err != nil { + if strings.Contains(err.Error(), "request body too large") { + return + } + http.Error(w, "Failed to read request body", http.StatusBadRequest) + return + } + + var req dto.OpenAITTSRequest + if err := json.Unmarshal(body, &req); err != nil { + http.Error(w, "Invalid JSON", http.StatusBadRequest) + return + } + + if req.Input == "" { + http.Error(w, "Input text is required", http.StatusBadRequest) + return + } + + if len(req.Input) > common.MaxTextLength { + http.Error(w, fmt.Sprintf("Input text too long (max %d characters)", common.MaxTextLength), http.StatusBadRequest) + return + } + + speed := req.Speed + if speed <= 0 { + speed = common.DefaultSpeed + } + if speed < common.MinSpeed { + speed = common.MinSpeed + } + if speed > common.MaxSpeed { + speed = common.MaxSpeed + } + + ttsStart := time.Now() + result, err := volcano.Synthesis(&setting.TTSConfig, volcanoClient, req.Input, speed) + duration := time.Since(ttsStart) + + if err != nil { + service.GlobalStats.AddRequest(false, duration, err.Error()) + http.Error(w, "TTS synthesis failed", http.StatusInternalServerError) + return + } + + service.GlobalStats.AddRequest(true, duration, "") + + w.Header().Set("Content-Type", "audio/wav") + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(result.AudioData))) + w.Header().Set("X-Request-Id", result.ReqID) + w.WriteHeader(http.StatusOK) + w.Write(result.AudioData) +} + +func HealthHandler(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + + if setting.TTSConfigErr != nil { + w.WriteHeader(http.StatusServiceUnavailable) + } else { + w.WriteHeader(http.StatusOK) + } + + totalRequests, successfulRequests, failedRequests, totalResponseTime, recentResponseTimes, lastErrors := service.GlobalStats.GetSnapshot() + + var errorRate float64 + if totalRequests > 0 { + errorRate = float64(failedRequests) / float64(totalRequests) * 100 + } + + var avgResponseTime float64 + if totalRequests > 0 { + avgResponseTime = totalResponseTime.Seconds() * 1000 / float64(totalRequests) + } + + envCheckStatus := setting.CheckEnvironmentVariables() + allEnvVarsSet := envCheckStatus["all_required_vars_set"].(bool) + + status := "ok" + if !allEnvVarsSet { + status = "configuration_error" + } + + response := dto.HealthResponse{ + Status: status, + Service: "ByteDance TTS to OpenAI API Adapter", + Version: "2.0.0 (v3 API)", + Uptime: fmt.Sprintf("%.0f seconds", time.Since(startTime).Seconds()), + StartTime: startTime.Format(time.RFC3339), + Memory: service.GetMemoryInfo(), + APIStats: dto.APIStatsResponse{ + TotalRequests: int(totalRequests), + SuccessfulRequests: successfulRequests, + FailedRequests: failedRequests, + ErrorRatePercent: fmt.Sprintf("%.2f", errorRate), + AvgResponseTimeMs: fmt.Sprintf("%.2f", avgResponseTime), + RecentResponseTimesMs: recentResponseTimes, + }, + Errors: dto.ErrorResponse{ + RecentErrorsCount: len(lastErrors), + }, + ConfigStatus: dto.ConfigStatusResponse{ + AllRequiredVarsSet: allEnvVarsSet, + ConfigError: setting.TTSConfigErr != nil, + }, + } + + json.NewEncoder(w).Encode(response) +} + +var startTime time.Time + +func SetStartTime(t time.Time) { + startTime = t +} diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..c40c000 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,21 @@ +services: + tts-api: + build: . + container_name: tts-api + ports: + - "${PORT:-8080}:8080" + environment: + - BYTEDANCE_TTS_API_KEY=${BYTEDANCE_TTS_API_KEY} + - BYTEDANCE_TTS_RESOURCE_ID=${BYTEDANCE_TTS_RESOURCE_ID} + - BYTEDANCE_TTS_SPEAKER=${BYTEDANCE_TTS_SPEAKER} + - BYTEDANCE_TTS_TIMEOUT=${BYTEDANCE_TTS_TIMEOUT:-30s} + - OPENAI_TTS_API_KEY=${OPENAI_TTS_API_KEY:-} + - ALLOWED_ORIGINS=${ALLOWED_ORIGINS:-} + - PORT=8080 + restart: unless-stopped + healthcheck: + test: ["CMD", "wget", "-qO-", "http://localhost:8080/health"] + interval: 30s + timeout: 5s + retries: 3 + start_period: 5s diff --git a/dto/health.go b/dto/health.go new file mode 100644 index 0000000..c1ba13b --- /dev/null +++ b/dto/health.go @@ -0,0 +1,31 @@ +package dto + +type HealthResponse struct { + Status string `json:"status"` + Service string `json:"service"` + Version string `json:"version"` + Uptime string `json:"uptime"` + StartTime string `json:"start_time"` + Memory map[string]interface{} `json:"memory"` + APIStats APIStatsResponse `json:"api_stats"` + Errors ErrorResponse `json:"errors"` + ConfigStatus ConfigStatusResponse `json:"config_status"` +} + +type APIStatsResponse struct { + TotalRequests int `json:"total_requests"` + SuccessfulRequests int64 `json:"successful_requests"` + FailedRequests int64 `json:"failed_requests"` + ErrorRatePercent string `json:"error_rate_percent"` + AvgResponseTimeMs string `json:"avg_response_time_ms"` + RecentResponseTimesMs []float64 `json:"recent_response_times_ms"` +} + +type ErrorResponse struct { + RecentErrorsCount int `json:"recent_errors_count"` +} + +type ConfigStatusResponse struct { + AllRequiredVarsSet bool `json:"all_required_vars_set"` + ConfigError bool `json:"config_error"` +} diff --git a/dto/tts.go b/dto/tts.go new file mode 100644 index 0000000..2e1a17f --- /dev/null +++ b/dto/tts.go @@ -0,0 +1,40 @@ +package dto + +import "time" + +type OpenAITTSRequest struct { + Model string `json:"model"` + Input string `json:"input"` + Voice string `json:"voice"` + ResponseFormat string `json:"response_format,omitempty"` + Speed float64 `json:"speed,omitempty"` +} + +type V3TTSResponse struct { + ReqID string `json:"reqid"` + Code int `json:"code"` + Message string `json:"message"` + Event string `json:"event"` + Sequence int `json:"sequence"` + Data string `json:"data"` + Sentence string `json:"sentence,omitempty"` + IsFinal bool `json:"is_final"` + Usage *V3Usage `json:"usage,omitempty"` +} + +type V3Usage struct { + TextWords int `json:"text_words"` +} + +type ByteDanceTTSConfig struct { + ApiKey string + ResourceId string + Speaker string + URL string + Timeout time.Duration +} + +type SynthesisResult struct { + AudioData []byte + ReqID string +} diff --git a/go.mod b/go.mod index 1533b9d..101db9f 100644 --- a/go.mod +++ b/go.mod @@ -1,4 +1,4 @@ -module bytedance-tts-openai-adapter +module github.com/volcano-tts/tts-api go 1.19 diff --git a/main.go b/main.go new file mode 100644 index 0000000..d759cbc --- /dev/null +++ b/main.go @@ -0,0 +1,82 @@ +package main + +import ( + "context" + "log" + "net/http" + "os" + "os/signal" + "syscall" + "time" + + "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/controller" + "github.com/volcano-tts/tts-api/middleware" + "github.com/volcano-tts/tts-api/router" + "github.com/volcano-tts/tts-api/service" + "github.com/volcano-tts/tts-api/setting" +) + +func main() { + log.SetFlags(log.LstdFlags | log.Lshortfile) + log.SetPrefix("[TTS-Server] ") + + middleware.InitAPIKeys() + middleware.InitCORSConfig() + middleware.InitRateLimiter() + setting.CheckStaticFiles() + service.InitStats() + controller.InitController() + + setting.TTSConfigErr = setting.InitTTSConfig() + if setting.TTSConfigErr != nil { + log.Printf("警告: 配置初始化失败: %v", setting.TTSConfigErr) + log.Printf("服务将继续运行,但TTS功能不可用,请检查环境变量配置") + } else { + log.Printf("配置初始化成功") + } + + controller.SetStartTime(time.Now()) + + r := router.Setup() + + port := os.Getenv("PORT") + if port == "" { + port = common.DefaultPort + } + + server := &http.Server{ + Addr: ":" + port, + Handler: r, + ReadTimeout: 30 * time.Second, + WriteTimeout: 120 * time.Second, + IdleTimeout: 60 * time.Second, + } + + quit := make(chan os.Signal, 1) + signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) + + go func() { + log.Printf("Starting ByteDance TTS to OpenAI API Adapter Server") + 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") + + if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed { + log.Fatalf("Server failed to start: %v", err) + } + }() + + <-quit + log.Println("Shutting down server...") + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + if err := server.Shutdown(ctx); err != nil { + log.Printf("Server forced to shutdown: %v", err) + } else { + log.Println("Server exited gracefully") + } +} diff --git a/middleware/auth.go b/middleware/auth.go new file mode 100644 index 0000000..a18f9a2 --- /dev/null +++ b/middleware/auth.go @@ -0,0 +1,59 @@ +package middleware + +import ( + "encoding/json" + "log" + "net/http" + "os" + "strings" +) + +var validAPIKeys []string + +func InitAPIKeys() { + apiKey := os.Getenv("OPENAI_TTS_API_KEY") + if apiKey != "" { + validAPIKeys = strings.Split(apiKey, ",") + for i, k := range validAPIKeys { + validAPIKeys[i] = strings.TrimSpace(k) + } + log.Printf("已配置 %d 个有效的API密钥", len(validAPIKeys)) + } else { + log.Println("警告: OPENAI_TTS_API_KEY 环境变量未设置,将允许所有请求") + } +} + +func ValidateAPIKey(r *http.Request) bool { + if len(validAPIKeys) == 0 { + return true + } + + authHeader := r.Header.Get("Authorization") + if authHeader == "" { + return false + } + + if !strings.HasPrefix(authHeader, "Bearer ") { + return false + } + + token := strings.TrimPrefix(authHeader, "Bearer ") + for _, validKey := range validAPIKeys { + if token == validKey { + return true + } + } + return false +} + +func SendJSONError(w http.ResponseWriter, statusCode int, message string, errType string, code string) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(statusCode) + json.NewEncoder(w).Encode(map[string]interface{}{ + "error": map[string]interface{}{ + "message": message, + "type": errType, + "code": code, + }, + }) +} diff --git a/middleware/cors.go b/middleware/cors.go new file mode 100644 index 0000000..22a0775 --- /dev/null +++ b/middleware/cors.go @@ -0,0 +1,115 @@ +package middleware + +import ( + "log" + "net/http" + "os" + "strings" +) + +var ( + allowedOrigins []string + allowAllOrigins bool + corsMaxAgeHeader = "86400" +) + +func normalizeOrigin(origin string) string { + origin = strings.TrimSpace(origin) + origin = strings.TrimRight(origin, "/") + return strings.ToLower(origin) +} + +func InitCORSConfig() { + origins := os.Getenv("ALLOWED_ORIGINS") + if origins == "" { + log.Println("警告: ALLOWED_ORIGINS 环境变量未设置") + log.Println("出于安全考虑,跨域请求将被拒绝。如需开放跨域请配置 ALLOWED_ORIGINS") + log.Println("开发环境可设置 ALLOWED_ORIGINS=* 允许所有来源(不可与凭据共用)") + return + } + + parts := strings.Split(origins, ",") + for _, p := range parts { + o := strings.TrimSpace(p) + if o == "" { + continue + } + if o == "*" { + allowAllOrigins = true + continue + } + allowedOrigins = append(allowedOrigins, normalizeOrigin(o)) + } + + if allowAllOrigins { + log.Println("警告: ALLOWED_ORIGINS=*,将允许所有来源跨域请求(不携带凭据)") + } + if len(allowedOrigins) > 0 { + log.Printf("已配置 %d 个允许的跨域来源白名单", len(allowedOrigins)) + } +} + +func isValidOrigin(origin string) bool { + if origin == "" || origin == "null" || origin == "nil" { + return false + } + if !strings.HasPrefix(origin, "http://") && !strings.HasPrefix(origin, "https://") { + return false + } + return true +} + +func matchOrigin(origin string) (string, bool) { + if !isValidOrigin(origin) { + return "", false + } + if allowAllOrigins { + return "*", true + } + normalized := normalizeOrigin(origin) + for _, allowed := range allowedOrigins { + if allowed == normalized { + return origin, true + } + } + return "", false +} + +func CORS(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + origin := r.Header.Get("Origin") + + 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") + } + } + } + + if r.Method == http.MethodOptions { + if origin != "" { + if _, matched := matchOrigin(origin); !matched { + w.WriteHeader(http.StatusNoContent) + return + } + } + w.WriteHeader(http.StatusNoContent) + return + } + + next.ServeHTTP(w, r) + }) +} diff --git a/middleware/logger.go b/middleware/logger.go new file mode 100644 index 0000000..6b80991 --- /dev/null +++ b/middleware/logger.go @@ -0,0 +1,35 @@ +package middleware + +import ( + "log" + "net/http" + "time" +) + +type statusRecorder struct { + http.ResponseWriter + statusCode int +} + +func (rec *statusRecorder) WriteHeader(code int) { + rec.statusCode = code + rec.ResponseWriter.WriteHeader(code) +} + +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) + duration := time.Since(start) + + log.Printf("%s %s %s %d %v", r.Method, r.RequestURI, r.RemoteAddr, rec.statusCode, duration) + }) +} diff --git a/middleware/ratelimit.go b/middleware/ratelimit.go new file mode 100644 index 0000000..6f22cf3 --- /dev/null +++ b/middleware/ratelimit.go @@ -0,0 +1,101 @@ +package middleware + +import ( + "net" + "net/http" + "strings" + "sync" + "time" + + "github.com/volcano-tts/tts-api/common" +) + +type RateLimiter struct { + requests map[string][]time.Time + mutex sync.Mutex + limit int + window time.Duration + lastCleanup time.Time +} + +var ( + GlobalRateLimiter *RateLimiter + ConcurrencySem chan struct{} +) + +func InitRateLimiter() { + GlobalRateLimiter = &RateLimiter{ + requests: make(map[string][]time.Time), + limit: common.RateLimitRequests, + window: common.RateLimitWindow, + } + ConcurrencySem = make(chan struct{}, common.MaxConcurrentRequests) +} + +func (rl *RateLimiter) Allow(key string) bool { + rl.mutex.Lock() + defer rl.mutex.Unlock() + + now := time.Now() + cutoff := now.Add(-rl.window) + + if now.Sub(rl.lastCleanup) > common.CleanupInterval { + rl.cleanup() + rl.lastCleanup = now + } + + timestamps := rl.requests[key] + valid := make([]time.Time, 0, len(timestamps)) + for _, ts := range timestamps { + if ts.After(cutoff) { + valid = append(valid, ts) + } + } + + if len(valid) >= rl.limit { + rl.requests[key] = valid + return false + } + + valid = append(valid, now) + rl.requests[key] = valid + return true +} + +func (rl *RateLimiter) cleanup() { + cutoff := time.Now().Add(-rl.window) + for k, v := range rl.requests { + valid := make([]time.Time, 0, len(v)) + for _, ts := range v { + if ts.After(cutoff) { + valid = append(valid, ts) + } + } + if len(valid) == 0 { + delete(rl.requests, k) + } else { + rl.requests[k] = valid + } + } +} + +func GetClientIP(r *http.Request) string { + xForwardedFor := r.Header.Get("X-Forwarded-For") + if xForwardedFor != "" { + ips := strings.Split(xForwardedFor, ",") + if len(ips) > 0 { + return strings.TrimSpace(ips[0]) + } + } + + xRealIP := r.Header.Get("X-Real-IP") + if xRealIP != "" { + return xRealIP + } + + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + return r.RemoteAddr + } + return host +} diff --git a/router/router.go b/router/router.go new file mode 100644 index 0000000..29b3731 --- /dev/null +++ b/router/router.go @@ -0,0 +1,27 @@ +package router + +import ( + "net/http" + + "github.com/gorilla/mux" + "github.com/volcano-tts/tts-api/controller" + "github.com/volcano-tts/tts-api/middleware" +) + +func Setup() *mux.Router { + r := mux.NewRouter() + + r.Use(middleware.CORS) + r.Use(middleware.Logger) + + r.HandleFunc("/v1/audio/speech", controller.OpenaiTTSHandler).Methods("POST", "OPTIONS") + r.HandleFunc("/health", controller.HealthHandler).Methods("GET") + r.HandleFunc("/dashboard", func(w http.ResponseWriter, r *http.Request) { + http.ServeFile(w, r, "health.html") + }).Methods("GET") + r.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/dashboard", http.StatusFound) + }).Methods("GET") + + return r +} diff --git a/service/stats.go b/service/stats.go new file mode 100644 index 0000000..3ac75a2 --- /dev/null +++ b/service/stats.go @@ -0,0 +1,90 @@ +package service + +import ( + "fmt" + "runtime" + "sync" + "time" + + "github.com/volcano-tts/tts-api/common" +) + +type Stats struct { + totalRequests int64 + successfulRequests int64 + failedRequests int64 + totalResponseTime time.Duration + recentResponseTimes []float64 + responseTimesIndex int + lastErrors []string + errorsIndex int + mutex sync.RWMutex +} + +var GlobalStats *Stats + +func InitStats() { + GlobalStats = &Stats{ + recentResponseTimes: make([]float64, common.MaxResponseTimes), + lastErrors: make([]string, common.MaxErrors), + } +} + +func (s *Stats) AddRequest(success bool, responseTime time.Duration, errMsg string) { + s.mutex.Lock() + defer s.mutex.Unlock() + + s.totalRequests++ + s.totalResponseTime += responseTime + + s.recentResponseTimes[s.responseTimesIndex] = responseTime.Seconds() * 1000 + s.responseTimesIndex = (s.responseTimesIndex + 1) % common.MaxResponseTimes + + 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 + s.errorsIndex = (s.errorsIndex + 1) % common.MaxErrors + } + } +} + +func (s *Stats) GetSnapshot() (totalRequests int64, successfulRequests int64, failedRequests int64, + totalResponseTime time.Duration, recentResponseTimes []float64, lastErrors []string) { + s.mutex.RLock() + defer s.mutex.RUnlock() + + totalRequests = s.totalRequests + successfulRequests = s.successfulRequests + failedRequests = s.failedRequests + totalResponseTime = s.totalResponseTime + + recentResponseTimes = make([]float64, 0, common.MaxResponseTimes) + for _, t := range s.recentResponseTimes { + if t > 0 { + recentResponseTimes = append(recentResponseTimes, t) + } + } + + lastErrors = make([]string, 0, common.MaxErrors) + for _, e := range s.lastErrors { + if e != "" { + lastErrors = append(lastErrors, e) + } + } + return +} + +func GetMemoryInfo() map[string]interface{} { + var m runtime.MemStats + runtime.ReadMemStats(&m) + return map[string]interface{}{ + "total_alloc": m.TotalAlloc, + "heap_alloc": m.HeapAlloc, + "heap_inuse": m.HeapInuse, + "goroutines": runtime.NumGoroutine(), + } +} diff --git a/setting/config.go b/setting/config.go new file mode 100644 index 0000000..5389f7b --- /dev/null +++ b/setting/config.go @@ -0,0 +1,91 @@ +package setting + +import ( + "fmt" + "log" + "os" + "time" + + "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/dto" +) + +var ( + TTSConfig dto.ByteDanceTTSConfig + TTSConfigErr error +) + +func InitTTSConfig() error { + apiKey := os.Getenv("BYTEDANCE_TTS_API_KEY") + resourceId := os.Getenv("BYTEDANCE_TTS_RESOURCE_ID") + speaker := os.Getenv("BYTEDANCE_TTS_SPEAKER") + + missingVars := []string{} + if apiKey == "" { + missingVars = append(missingVars, "BYTEDANCE_TTS_API_KEY") + } + if resourceId == "" { + missingVars = append(missingVars, "BYTEDANCE_TTS_RESOURCE_ID") + } + if speaker == "" { + missingVars = append(missingVars, "BYTEDANCE_TTS_SPEAKER") + } + + if len(missingVars) > 0 { + return fmt.Errorf("缺少必需的环境变量: %v", missingVars) + } + + url := "https://openspeech.bytedance.com/api/v3/tts/unidirectional" + + timeout := common.DefaultTimeout + if timeoutStr := os.Getenv("BYTEDANCE_TTS_TIMEOUT"); timeoutStr != "" { + if parsedTimeout, err := time.ParseDuration(timeoutStr); err == nil { + timeout = parsedTimeout + } else { + log.Printf("无效的超时设置 '%s',使用默认值: %v", timeoutStr, timeout) + } + } + + TTSConfig = dto.ByteDanceTTSConfig{ + ApiKey: apiKey, + ResourceId: resourceId, + Speaker: speaker, + URL: url, + Timeout: timeout, + } + return nil +} + +func CheckEnvironmentVariables() map[string]interface{} { + requiredVars := map[string]bool{ + "BYTEDANCE_TTS_API_KEY": os.Getenv("BYTEDANCE_TTS_API_KEY") != "", + "BYTEDANCE_TTS_RESOURCE_ID": os.Getenv("BYTEDANCE_TTS_RESOURCE_ID") != "", + "BYTEDANCE_TTS_SPEAKER": os.Getenv("BYTEDANCE_TTS_SPEAKER") != "", + } + + missingVars := []string{} + for varName, isSet := range requiredVars { + if !isSet { + missingVars = append(missingVars, varName) + } + } + + optionalVars := map[string]bool{ + "BYTEDANCE_TTS_TIMEOUT": os.Getenv("BYTEDANCE_TTS_TIMEOUT") != "", + "OPENAI_TTS_API_KEY": os.Getenv("OPENAI_TTS_API_KEY") != "", + "PORT": os.Getenv("PORT") != "", + } + + return map[string]interface{}{ + "all_required_vars_set": len(missingVars) == 0, + "missing_required_vars": missingVars, + "required_vars_set": requiredVars, + "optional_vars_set": optionalVars, + } +} + +func CheckStaticFiles() { + if _, err := os.Stat("health.html"); os.IsNotExist(err) { + log.Println("警告: health.html 不存在,/dashboard 路由将返回 404") + } +} diff --git a/tts-api.exe b/tts-api.exe new file mode 100644 index 0000000..15738ef Binary files /dev/null and b/tts-api.exe differ diff --git a/tts_api_architecture.html b/tts_api_architecture.html new file mode 100644 index 0000000..76a0607 --- /dev/null +++ b/tts_api_architecture.html @@ -0,0 +1,444 @@ + + + + + +TTS-API 架构设计 — TTS 版 New-API + + + + +

TTS-API

+

TTS 版 New-API 架构设计 · 多 Provider 统一 TTS 网关 · 2026-05-21

+ +

1. 项目定位

+ +

参考 new-api 的设计理念,TTS-API 定位为企业级 TTS 统一网关与资产管理平台,核心能力:

+ + + + + + + +
能力维度说明
统一接入以 OpenAI /v1/audio/speech 为唯一入口,屏蔽火山引擎、阿里、Azure、讯飞等上游差异
统一音色定义标准音色命名体系,自动映射到各 Provider 的实际音色 ID
统一计费按字符数 / 时长计费,支持配额管理与成本核算
统一治理权限分组、速率限制、渠道故障切换、审计日志、可视化看板
+ +

2. 与 new-api 的关键差异

+ + + + + + + + + + +
维度new-api(LLM)TTS-API(本项目)
核心接口/v1/chat/completions/v1/audio/speech
协议转换OpenAI ↔ Claude ↔ Gemini 互转只需转到各 Provider 原生格式(单向)
模型映射模型名 → 渠道选择音色映射:标准 voice → 各 Provider 实际音色 ID
(这是最大难点)
输出处理文本 / JSON 流二进制音频流,需处理格式转换(wav/mp3/pcm)
计费单位Token 数字符数 + 音频时长
缓存语义缓存(相同问题命中)音频缓存(相同文本+音色 → 直接返回已合成音频)
流式SSE 文本流音频流式推送(边合成边返回,首字节延迟是关键指标)
+ +

3. 整体分层架构

+ +
+┌─────────────────────────────────────────────────────────────────┐ +│ 客户端 / 应用层 │ +│ OpenAI SDK │ REST API │ Web 管理后台 │ +└────────────────────────────┬────────────────────────────────────┘ + │ +┌────────────────────────────▼────────────────────────────────────┐ +│ Gin HTTP Server │ +│ ┌──────────────────────────────────────────────────────────────┐│ +│ │ 路由层 (router/) ││ +│ │ /v1/audio/speech /api/* (管理) /web/* (前端静态) ││ +│ └──────────────────────────────────────────────────────────────┘│ +└────────────────────────────┬────────────────────────────────────┘ + │ +┌────────────────────────────▼────────────────────────────────────┐ +│ 中间件层 (middleware/) │ +│ 认证(JWT/API Key) │ 限流(IP/用户) │ 日志 │ CORS │ 请求分发 │ +└────────────────────────────┬────────────────────────────────────┘ + │ +┌────────────────────────────▼────────────────────────────────────┐ +│ 控制器层 (controller/) │ +│ ┌──────────┐ ┌──────────┐ ┌──────────┐ ┌──────────────┐ │ +│ │ 语音合成 │ │ 用户管理 │ │ 渠道管理 │ │ 计费 & 统计 │ │ +│ │ controller│ │controller │ │controller │ │ controller │ │ +│ └──────────┘ └──────────┘ └──────────┘ └──────────────┘ │ +└────────────────────────────┬────────────────────────────────────┘ + │ + ┌─────────────────────┼─────────────────────┐ + │ │ │ +┌──────▼──────┐ ┌────────▼────────┐ ┌───────▼──────┐ +│ 服务层 │ │ 适配器层 │ │ 数据层 │ +│ (service/) │ │ (adapter/) │ │ (model/) │ +│ │ │ │ │ │ +│ · 配额管理 │ │ 接口定义 │ │ GORM ORM │ +│ · 音频缓存 │ │ · volcano │ │ │ +│ · 音色映射 │ │ · aliyun │ │ │ +│ · 格式转换 │ │ · azure │ │ │ +│ · 计费服务 │ │ · tencent │ │ │ +│ · 渠道调度 │ │ · xunfei │ │ │ +└──────────────┘ │ · openai │ └───────┬──────┘ + │ · fish_audio │ │ + │ · bert_vits │ ┌───────┼───────┐ + └─────────────────┘ │ │ │ + ┌──────▼──┐ ┌──▼──┐ ┌─▼────┐ + │ SQLite │ │MySQL│ │ PG │ + └─────────┘ └─────┘ └──────┘ +
+ +

4. 请求处理全流程

+ +
+ 客户端发送 OpenAI 格式 TTS 请求 + │ + ▼ + [1] Gin 路由匹配 → /v1/audio/speech + │ + ▼ + [2] 中间件链:认证(API Key) → 用户级限流 → 请求日志 + │ + ▼ + [3] 控制器:解析请求 { model, input, voice, speed, response_format } + │ + ▼ + [4] 音色映射服务:标准 voice 名 → 查找可用渠道 → 映射为渠道实际音色ID + │ 例: "gentle_male" → 火山引擎(zh_male_qingxin) + │ → Azure(zh-CN-YunxiNeural) + │ → 阿里(cosyvoice-v1-longxiaochun) + ▼ + [5] 渠道调度器:按权重+可用性选择最优渠道,失败自动切换 + │ + ▼ + [6] 适配器:将 OpenAI 请求转为 Provider 原生格式,发送请求 + │ + ▼ + [7] 音频处理:接收二进制音频 → 格式转换(如需) → 写入音频缓存 + │ + ▼ + [8] 计费结算:按字符数/时长扣费,记录日志 + │ + ▼ + [9] 返回响应:Content-Type: audio/wav,流式或整段返回 +
+ +

5. 核心设计:适配器接口

+ +

参考 new-api 的 Adaptor 接口设计,TTS 版适配器接口如下:

+ +
+// adapter.go — TTS Provider 统一接口 + +type TTSAdapter interface { + // 初始化:传入渠道配置(API Key / Resource ID / 默认音色等) + Init(info *TTSRelayInfo) error + + // 构建上游请求 URL + BuildRequestURL(info *TTSRelayInfo) (string, error) + + // 设置请求头(鉴权、Content-Type 等) + SetupRequestHeader(c *gin.Context, req *http.Request, info *TTSRelayInfo) error + + // 核心:将 OpenAI 格式请求转为 Provider 原生请求体 + ConvertRequest(info *TTSRelayInfo, req *dto.OpenAITTSRequest) (any, error) + + // 发送请求到上游 + DoRequest(c *gin.Context, info *TTSRelayInfo, body io.Reader) (*http.Response, error) + + // 处理上游响应:提取音频、统计字符数/时长、返回标准结构 + DoResponse(c *gin.Context, resp *http.Response, info *TTSRelayInfo) (*dto.TTSUsage, *dto.TTSError) + + // 返回此渠道支持的音色列表(用于音色映射表构建) + GetVoiceList() []ProviderVoice + + // 渠道标识 + GetChannelName() string + + // 是否支持流式 TTS + SupportStreaming() bool +} + +// ProviderVoice 各 Provider 的音色结构 +type ProviderVoice struct { + ProviderID string // Provider 内部音色 ID,如 "zh_female_qingxin" + Language string // zh-CN / en-US / ja-JP + Gender string // male / female + Style string // 风格标签,如 "news" / "story" / "chat" + Description string // 音色描述 +} +
+ +

6. 核心设计:音色映射系统

+ +

这是 TTS 网关区别于 LLM 网关的最大难点与核心创新点。LLM 网关只需按模型名路由,但 TTS 需要一套跨 Provider 的音色统一体系。

+ +
+ 标准音色命名空间 + ┌──────────────────────────────────┐ + │ tts-1-gentle-male │ + │ tts-1-gentle-female │ + │ tts-1-news-male │ + │ tts-1-story-female │ + │ tts-1-casual-male │ + │ ... │ + └──────────┬───────────────────────┘ + │ 音色映射表 (voice_mapping) + │ + ┌───────────────┼───────────────┬───────────────┐ + │ │ │ │ + ▼ ▼ ▼ ▼ +┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐ +│ 火山引擎 │ │ Azure │ │ 阿里 │ │ 讯飞 │ +│ │ │ │ │ │ │ │ +│zh_male_ │ │zh-CN- │ │cosyvoice│ │x4_ling │ +│qingxin │ │Yunxi │ │-v1-long │ │xiaoxuan │ +│ │ │Neural │ │xiaochun │ │ │ +└─────────┘ └─────────┘ └─────────┘ └─────────┘ +
+ +

6.1 音色映射表结构

+ +
+// voice_mapping 表 (数据库) +type VoiceMapping struct { + ID uint + StandardVoice string // "tts-1-gentle-male" + ChannelID uint // 渠道 ID + ProviderVoice string // Provider 原始音色 ID + Priority int // 优先级(同一标准音色多渠道时,优先选谁) + IsDefault bool // 是否为该标准音色的默认渠道 +} + +// 示例数据: +// standard_voice | channel | provider_voice | priority +// tts-1-gentle-male | 火山引擎 | zh_male_qingxin | 1 +// tts-1-gentle-male | Azure | zh-CN-YunxiNeural | 2 +// tts-1-gentle-male | 阿里 | cosyvoice-v1-longxiaochun| 3 +
+ +

6.2 音色发现与自动映射建议

+ +

每个 Provider 适配器实现 GetVoiceList(),系统启动或渠道更新时自动拉取,通过语言+性别+风格标签与标准命名空间做相似度匹配,自动生成映射建议,管理员在 Web UI 审核确认即可。

+ +

7. 渠道调度与故障切换

+ +
+ 请求 → 查找"tts-1-gentle-male"可用的渠道列表 + │ + ▼ + 按 priority 排序 → 加权随机选择一个渠道 + │ + ▼ + 适配器发送请求 → 成功? + ├── 是 → 返回音频,记录成功 + │ + └── 否 → 标记渠道失败 + │ + ▼ + 自动切换 priority+1 的渠道重试 + │ + ▼ + 所有渠道失败 → 返回 503 + 错误详情 +
+ +

8. 项目目录结构

+ +
+tts-api/ +├── main.go # 入口:初始化 DB、路由、启动服务 +├── go.mod / go.sum +├── .env.example # 环境变量示例 +├── Dockerfile # 多阶段构建 +├── docker-compose.yml +│ +├── router/ # 路由层 +│ ├── main.go # 路由聚合 +│ ├── api-router.go # /api/* 管理接口 +│ ├── relay-router.go # /v1/audio/speech TTS 代理 +│ └── web-router.go # /web/* 管理后台静态资源 +│ +├── middleware/ # 中间件 +│ ├── auth.go # JWT + API Key 认证 +│ ├── rate-limit.go # 用户/IP 级别限流 +│ ├── cors.go # 跨域 +│ └── logger.go # 请求日志 +│ +├── controller/ # 控制器 +│ ├── tts.go # TTS 合成入口(核心) +│ ├── channel.go # 渠道 CRUD +│ ├── user.go # 用户管理 +│ ├── token.go # API Key 管理 +│ ├── voice.go # 音色映射管理 +│ └── billing.go # 计费统计 +│ +├── service/ # 服务层 +│ ├── voice_mapping/ # 音色映射服务(核心) +│ │ └── matcher.go # 自动匹配 & 建议 +│ ├── channel_scheduler/ # 渠道调度(权重/故障切换) +│ │ └── scheduler.go +│ ├── audio_cache/ # 音频缓存(相同文本+音色命中) +│ │ └── cache.go +│ ├── audio_convert/ # 音频格式转换(ffmpeg 封装) +│ │ └── converter.go +│ ├── billing/ # 计费服务 +│ │ └── billing.go +│ └── quota/ # 配额管理 +│ └── quota.go +│ +├── adapter/ # 适配器层(核心!) +│ ├── adapter.go # TTSAdapter 接口定义 +│ ├── volcano/ # 火山引擎 TTS +│ │ └── volcano.go +│ ├── aliyun/ # 阿里云 CosyVoice / 百炼 +│ │ └── aliyun.go +│ ├── azure/ # 微软 Azure TTS +│ │ └── azure.go +│ ├── tencent/ # 腾讯云 TTS +│ │ └── tencent.go +│ ├── xunfei/ # 讯飞 TTS +│ │ └── xunfei.go +│ ├── openai/ # OpenAI TTS(基准) +│ │ └── openai.go +│ ├── fish_audio/ # Fish Audio(开源) +│ │ └── fish_audio.go +│ └── bert_vits/ # Bert-VITS2(开源自建) +│ └── bert_vits.go +│ +├── model/ # 数据模型 (GORM) +│ ├── user.go +│ ├── channel.go # 渠道(Provider 配置) +│ ├── token.go # API Key / Token +│ ├── voice_mapping.go # 音色映射 +│ ├── usage_record.go # 用量记录 +│ └── audio_cache.go # 音频缓存记录 +│ +├── dto/ # 请求/响应结构体 +│ ├── openai_tts.go # OpenAI TTS 请求/响应格式 +│ ├── relay_info.go # 中继上下文(TTSRelayInfo) +│ └── common.go # 通用响应 +│ +├── setting/ # 配置管理 (Viper) +│ └── setting.go +│ +├── common/ # 通用工具 +│ ├── utils.go +│ └── constants.go +│ +└── web/ # React 管理后台 + ├── src/ + │ ├── pages/ + │ │ ├── Dashboard # 数据看板 + │ │ ├── Channels # 渠道管理 + │ │ ├── VoiceMapping # 音色映射配置 + │ │ ├── Users # 用户管理 + │ │ ├── Tokens # API Key + │ │ ├── Billing # 计费 & 用量 + │ │ └── Logs # 调用日志 + │ └── ... + └── package.json +
+ +

9. 管理后台页面规划

+ + + + + + + + + + +
页面功能
Dashboard今日合成次数、字符数、时长、费用、渠道健康状态、QPS 曲线
渠道管理添加/编辑 Provider(API Key、Resource ID、权重、并发上限)
音色映射核心页面:标准音色 ↔ 各渠道音色 ID 的映射表,支持自动匹配建议与手动调整
API Key生成/管理用户 API Key,绑定分组与配额
用户管理用户 CRUD、分组、角色(Admin / User)
计费统计按用户/渠道/日期维度的用量与费用报表
调用日志每次 TTS 请求的详细日志(文本、音色、渠道、耗时、费用)
+ +

10. 一期 vs 二期路线图

+ +

一期(MVP,你现有的 Volcano-Engine-TTS-UI 升级版)

+ + + + + + + + +
模块内容
适配器火山引擎 + Azure + 阿里云,3 个 Provider
接口/v1/audio/speech,OpenAI 兼容
音色映射硬编码映射表(配置文件),先跑通再抽象
计费简单字符数计数 + 日志
管理后台极简版:渠道配置页面 + 用量看板
数据库SQLite,单文件部署
+ +

二期(完整版)

+ + + + + + + + + + +
模块内容
适配器扩到 8+ Provider(讯飞、腾讯、Fish Audio、Bert-VITS2、OpenAI)
音色映射数据库驱动 + Web UI 可视化管理 + 自动匹配建议
计费字符数/时长双维度计费,用户配额,欠费阻断
缓存Redis 音频缓存,相同文本+音色直接命中
流式支持流式 TTS(SSE 推送音频 chunk),降低首字节延迟
管理后台完整 React 后台(参考 new-api 的 Semi Design UI)
数据库MySQL / PostgreSQL 支持
部署Docker Compose 一键部署
+ +

11. 关键工程建议

+ +
    +
  1. 从你现有的火山引擎适配器起步,先抽象出 TTSAdapter 接口,接入 2~3 个 Provider 验证接口设计是否合理,不要一上来就搞 8 个适配器。
  2. +
  3. 音色映射先硬编码,跑通流程后再做成数据库驱动的 Web UI。映射表是长期维护工作,需要社区共建。
  4. +
  5. 音频格式转换用 ffmpeg,Go 侧通过 exec.Command 调用或使用 go-ffmpeg 绑定。各 Provider 输出格式不同(wav / mp3 / pcm),统一转码是刚需。
  6. +
  7. 渠道调度直接复用 new-api 的加权随机 + 故障重试思路,这是成熟的模式,不需要重新发明。
  8. +
  9. 管理后台前期可以不写,SQLite + 配置文件就能用;等 Provider 多了再补 React 前端。
  10. +
  11. 考虑直接 fork new-api 改造:new-api 的渠道管理、用户系统、计费框架、中间件、部署方案都是现成的,你只需要把 relay 层的 LLM 适配器替换成 TTS 适配器,再加音色映射模块。这比从零搭建快得多。
  12. +
+ + + + + \ No newline at end of file diff --git a/tts_server.go b/tts_server.go deleted file mode 100644 index 5f5f759..0000000 --- a/tts_server.go +++ /dev/null @@ -1,882 +0,0 @@ -package main - -import ( - "bufio" - "bytes" - "context" - "encoding/base64" - "encoding/json" - "fmt" - "io" - "log" - "net" - "net/http" - "os" - "os/signal" - "runtime" - "strings" - "sync" - "syscall" - "time" - - "github.com/google/uuid" - "github.com/gorilla/mux" -) - -const ( - DEFAULT_PORT = "8080" - DEFAULT_TIMEOUT = 30 * time.Second - MAX_TEXT_LENGTH = 5000 - MIN_SPEED = 0.25 - MAX_SPEED = 4.0 - DEFAULT_SPEED = 1.0 - MAX_REQUEST_BODY_SIZE = 1024 * 1024 - RATE_LIMIT_REQUESTS = 100 - RATE_LIMIT_WINDOW = time.Minute - MAX_RESPONSE_TIMES = 100 - MAX_ERRORS = 10 - MAX_CONCURRENT_REQUESTS = 10 -) - -type V3TTSResponse struct { - ReqID string `json:"reqid"` - Code int `json:"code"` - Message string `json:"message"` - Event string `json:"event"` - Sequence int `json:"sequence"` - Data string `json:"data"` - Sentence string `json:"sentence,omitempty"` - IsFinal bool `json:"is_final"` - Usage *Usage `json:"usage,omitempty"` -} - -type Usage struct { - TextWords int `json:"text_words"` -} - -type OpenAITTSRequest struct { - Model string `json:"model"` - Input string `json:"input"` - Voice string `json:"voice"` - ResponseFormat string `json:"response_format,omitempty"` - Speed float64 `json:"speed,omitempty"` -} - -type ByteDanceTTSConfig struct { - ApiKey string - ResourceId string - Speaker string - URL string - Timeout time.Duration -} - -type RateLimiter struct { - requests map[string][]time.Time - mutex sync.Mutex - limit int - window time.Duration - lastCleanup time.Time -} - -const cleanupInterval = time.Hour - -type Stats struct { - totalRequests int64 - successfulRequests int64 - failedRequests int64 - totalResponseTime time.Duration - recentResponseTimes []float64 - responseTimesIndex int - lastErrors []string - errorsIndex int - mutex sync.RWMutex -} - -var ( - VALID_API_KEYS []string - ttsConfig ByteDanceTTSConfig - ttsConfigErr error - globalHTTPClient *http.Client - apiStats *Stats - rateLimiter *RateLimiter - concurrencySem chan struct{} -) - -func init() { - globalHTTPClient = &http.Client{ - Timeout: DEFAULT_TIMEOUT, - Transport: &http.Transport{ - MaxIdleConns: 100, - MaxIdleConnsPerHost: 20, - IdleConnTimeout: 90 * time.Second, - TLSHandshakeTimeout: 10 * time.Second, - }, - } - - apiStats = &Stats{ - recentResponseTimes: make([]float64, MAX_RESPONSE_TIMES), - lastErrors: make([]string, MAX_ERRORS), - } - - rateLimiter = &RateLimiter{ - requests: make(map[string][]time.Time), - limit: RATE_LIMIT_REQUESTS, - window: RATE_LIMIT_WINDOW, - } - - concurrencySem = make(chan struct{}, MAX_CONCURRENT_REQUESTS) -} - -func (rl *RateLimiter) Allow(key string) bool { - rl.mutex.Lock() - defer rl.mutex.Unlock() - - now := time.Now() - cutoff := now.Add(-rl.window) - - if now.Sub(rl.lastCleanup) > cleanupInterval { - rl.cleanup() - rl.lastCleanup = now - } - - timestamps := rl.requests[key] - valid := make([]time.Time, 0, len(timestamps)) - for _, ts := range timestamps { - if ts.After(cutoff) { - valid = append(valid, ts) - } - } - - if len(valid) >= rl.limit { - rl.requests[key] = valid - return false - } - - valid = append(valid, now) - rl.requests[key] = valid - return true -} - -func (rl *RateLimiter) cleanup() { - cutoff := time.Now().Add(-rl.window) - for k, v := range rl.requests { - valid := make([]time.Time, 0, len(v)) - for _, ts := range v { - if ts.After(cutoff) { - valid = append(valid, ts) - } - } - if len(valid) == 0 { - delete(rl.requests, k) - } else { - rl.requests[k] = valid - } - } -} - -func initTTSConfig() error { - apiKey := os.Getenv("BYTEDANCE_TTS_API_KEY") - resourceId := os.Getenv("BYTEDANCE_TTS_RESOURCE_ID") - speaker := os.Getenv("BYTEDANCE_TTS_SPEAKER") - - missingVars := []string{} - - if apiKey == "" { - missingVars = append(missingVars, "BYTEDANCE_TTS_API_KEY") - } - if resourceId == "" { - missingVars = append(missingVars, "BYTEDANCE_TTS_RESOURCE_ID") - } - if speaker == "" { - missingVars = append(missingVars, "BYTEDANCE_TTS_SPEAKER") - } - - if len(missingVars) > 0 { - return fmt.Errorf("缺少必需的环境变量: %v", missingVars) - } - - url := "https://openspeech.bytedance.com/api/v3/tts/unidirectional" - - timeout := DEFAULT_TIMEOUT - if timeoutStr := os.Getenv("BYTEDANCE_TTS_TIMEOUT"); timeoutStr != "" { - if parsedTimeout, err := time.ParseDuration(timeoutStr); err == nil { - timeout = parsedTimeout - } else { - log.Printf("无效的超时设置 '%s',使用默认值: %v", timeoutStr, timeout) - } - } - - ttsConfig = ByteDanceTTSConfig{ - ApiKey: apiKey, - ResourceId: resourceId, - Speaker: speaker, - URL: url, - Timeout: timeout, - } - - return nil -} - -func initAPIKeys() { - apiKey := os.Getenv("OPENAI_TTS_API_KEY") - if apiKey != "" { - VALID_API_KEYS = strings.Split(apiKey, ",") - for i, k := range VALID_API_KEYS { - VALID_API_KEYS[i] = strings.TrimSpace(k) - } - log.Printf("已配置 %d 个有效的API密钥", len(VALID_API_KEYS)) - } else { - log.Println("警告: OPENAI_TTS_API_KEY 环境变量未设置,将允许所有请求") - } -} - -func checkEnvironmentVariables() map[string]interface{} { - requiredVars := map[string]bool{ - "BYTEDANCE_TTS_API_KEY": os.Getenv("BYTEDANCE_TTS_API_KEY") != "", - "BYTEDANCE_TTS_RESOURCE_ID": os.Getenv("BYTEDANCE_TTS_RESOURCE_ID") != "", - "BYTEDANCE_TTS_SPEAKER": os.Getenv("BYTEDANCE_TTS_SPEAKER") != "", - } - - missingVars := []string{} - for varName, isSet := range requiredVars { - if !isSet { - missingVars = append(missingVars, varName) - } - } - - optionalVars := map[string]bool{ - "BYTEDANCE_TTS_TIMEOUT": os.Getenv("BYTEDANCE_TTS_TIMEOUT") != "", - "OPENAI_TTS_API_KEY": os.Getenv("OPENAI_TTS_API_KEY") != "", - "PORT": os.Getenv("PORT") != "", - } - - return map[string]interface{}{ - "all_required_vars_set": len(missingVars) == 0, - "missing_required_vars": missingVars, - "required_vars_set": requiredVars, - "optional_vars_set": optionalVars, - } -} - -func httpPostStream(url string, headers map[string]string, body []byte, timeout time.Duration) (*http.Response, error) { - req, err := http.NewRequest(http.MethodPost, url, bytes.NewBuffer(body)) - if err != nil { - return nil, err - } - for key, value := range headers { - req.Header.Set(key, value) - } - - ctx, cancel := context.WithTimeout(context.Background(), timeout) - defer cancel() - - req = req.WithContext(ctx) - - return globalHTTPClient.Do(req) -} - -func convertSpeedToSpeechRate(speed float64) int { - if speed <= 0.5 { - return -50 - } - if speed >= 2.0 { - return 100 - } - return int((speed - 1.0) * 100) -} - -type SynthesisResult struct { - AudioData []byte - ReqID string -} - -func synthesis(text string, speed float64) (*SynthesisResult, error) { - reqID := uuid.NewString() - - speechRate := convertSpeedToSpeechRate(speed) - - params := map[string]interface{}{ - "user": map[string]interface{}{ - "uid": "uid", - }, - "namespace": "BidirectionalTTS", - "req_params": map[string]interface{}{ - "text": text, - "speaker": ttsConfig.Speaker, - "audio_params": map[string]interface{}{ - "format": "wav", - "sample_rate": 24000, - "speech_rate": speechRate, - }, - }, - } - - headers := map[string]string{ - "Content-Type": "application/json", - "Connection": "keep-alive", - "X-Api-Resource-Id": ttsConfig.ResourceId, - "X-Api-Request-Id": reqID, - "X-Api-Key": ttsConfig.ApiKey, - } - - bodyStr, err := json.Marshal(params) - if err != nil { - log.Printf("JSON marshal fail: %v", err) - return nil, err - } - - resp, err := httpPostStream(ttsConfig.URL, headers, bodyStr, ttsConfig.Timeout) - if err != nil { - log.Printf("http post fail: %v", err) - return nil, err - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - 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) - } - - var audioData []byte - scanner := bufio.NewScanner(resp.Body) - scanner.Buffer(make([]byte, 1024*1024), 8*1024*1024) - - for scanner.Scan() { - line := scanner.Bytes() - if len(line) == 0 { - continue - } - - var v3Resp V3TTSResponse - if err := json.Unmarshal(line, &v3Resp); err != nil { - log.Printf("unmarshal chunk fail: %v, line: %s", err, string(line)) - continue - } - - if v3Resp.Code == 20000000 { - if v3Resp.Usage != nil { - log.Printf("TTS synthesis completed, usage: %+v", v3Resp.Usage) - } - for scanner.Scan() { - } - break - } - - if v3Resp.Code != 0 { - log.Printf("TTS service error: code=%d, message=%s", v3Resp.Code, v3Resp.Message) - return nil, fmt.Errorf("TTS service error: %s", v3Resp.Message) - } - - if v3Resp.Data != "" { - chunk, err := base64.StdEncoding.DecodeString(v3Resp.Data) - if err != nil { - log.Printf("base64 decode fail: %v", err) - return nil, err - } - audioData = append(audioData, chunk...) - } else if v3Resp.Sentence != "" { - log.Printf("Received sentence info (sequence %d): %s", v3Resp.Sequence, v3Resp.Sentence) - } - } - - if err := scanner.Err(); err != nil { - log.Printf("read stream fail: %v", err) - return nil, err - } - - if len(audioData) == 0 { - return nil, fmt.Errorf("no audio data received") - } - - return &SynthesisResult{AudioData: audioData, ReqID: reqID}, nil -} - -func validateAPIKey(r *http.Request) bool { - if len(VALID_API_KEYS) == 0 { - return true - } - - authHeader := r.Header.Get("Authorization") - if authHeader == "" { - return false - } - - if !strings.HasPrefix(authHeader, "Bearer ") { - return false - } - - token := strings.TrimPrefix(authHeader, "Bearer ") - for _, validKey := range VALID_API_KEYS { - if token == validKey { - return true - } - } - return false -} - -func getClientIP(r *http.Request) string { - xForwardedFor := r.Header.Get("X-Forwarded-For") - if xForwardedFor != "" { - ips := strings.Split(xForwardedFor, ",") - if len(ips) > 0 { - return strings.TrimSpace(ips[0]) - } - } - - xRealIP := r.Header.Get("X-Real-IP") - if xRealIP != "" { - return xRealIP - } - - host, _, err := net.SplitHostPort(r.RemoteAddr) - if err != nil { - return r.RemoteAddr - } - return host -} - -func openaiTTSHandler(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost { - http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) - return - } - - if !validateAPIKey(r) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusUnauthorized) - json.NewEncoder(w).Encode(map[string]interface{}{ - "error": map[string]interface{}{ - "message": "Invalid API key provided.", - "type": "invalid_request_error", - "code": "invalid_api_key", - }, - }) - return - } - - if ttsConfigErr != nil { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusServiceUnavailable) - json.NewEncoder(w).Encode(map[string]interface{}{ - "error": map[string]interface{}{ - "message": fmt.Sprintf("TTS service configuration error: %v. Please check environment variables and restart the service.", ttsConfigErr), - "type": "configuration_error", - "code": "service_unavailable", - }, - }) - return - } - - select { - case concurrencySem <- struct{}{}: - defer func() { <-concurrencySem }() - default: - log.Printf("警告: 已达到最大并发请求数限制,拒绝请求 - 客户端IP: %s", getClientIP(r)) - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusServiceUnavailable) - json.NewEncoder(w).Encode(map[string]interface{}{ - "error": map[string]interface{}{ - "message": "Server is busy, maximum concurrent requests reached. Please try again later.", - "type": "concurrency_limit_error", - "code": "max_concurrent_requests", - }, - }) - return - } - - clientIP := getClientIP(r) - if !rateLimiter.Allow(clientIP) { - log.Printf("警告: 已超过IP速率限制,拒绝请求 - 客户端IP: %s", clientIP) - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusTooManyRequests) - json.NewEncoder(w).Encode(map[string]interface{}{ - "error": map[string]interface{}{ - "message": "Rate limit exceeded. Please try again later.", - "type": "rate_limit_error", - "code": "rate_limit_exceeded", - }, - }) - return - } - - r.Body = http.MaxBytesReader(w, r.Body, MAX_REQUEST_BODY_SIZE) - body, err := io.ReadAll(r.Body) - if err != nil { - if strings.Contains(err.Error(), "request body too large") { - return - } - http.Error(w, "Failed to read request body", http.StatusBadRequest) - return - } - - var req OpenAITTSRequest - if err := json.Unmarshal(body, &req); err != nil { - http.Error(w, "Invalid JSON", http.StatusBadRequest) - return - } - - if req.Input == "" { - http.Error(w, "Input text is required", http.StatusBadRequest) - return - } - - if len(req.Input) > MAX_TEXT_LENGTH { - http.Error(w, fmt.Sprintf("Input text too long (max %d characters)", MAX_TEXT_LENGTH), http.StatusBadRequest) - return - } - - speed := req.Speed - if speed <= 0 { - speed = DEFAULT_SPEED - } - if speed < MIN_SPEED { - speed = MIN_SPEED - } - if speed > MAX_SPEED { - speed = MAX_SPEED - } - - ttsStart := time.Now() - result, err := synthesis(req.Input, speed) - duration := time.Since(ttsStart) - - if err != nil { - addRequestStats(false, duration, err.Error()) - http.Error(w, "TTS synthesis failed", http.StatusInternalServerError) - return - } - - addRequestStats(true, duration, "") - - w.Header().Set("Content-Type", "audio/wav") - w.Header().Set("Content-Length", fmt.Sprintf("%d", len(result.AudioData))) - w.Header().Set("X-Request-Id", result.ReqID) - w.WriteHeader(http.StatusOK) - w.Write(result.AudioData) -} - -func addRequestStats(success bool, responseTime time.Duration, errMsg string) { - apiStats.mutex.Lock() - defer apiStats.mutex.Unlock() - - apiStats.totalRequests++ - apiStats.totalResponseTime += responseTime - - apiStats.recentResponseTimes[apiStats.responseTimesIndex] = responseTime.Seconds() * 1000 - apiStats.responseTimesIndex = (apiStats.responseTimesIndex + 1) % MAX_RESPONSE_TIMES - - if success { - apiStats.successfulRequests++ - } else { - apiStats.failedRequests++ - if errMsg != "" { - errInfo := fmt.Sprintf("%s: %s", time.Now().Format(time.RFC3339), errMsg) - apiStats.lastErrors[apiStats.errorsIndex] = errInfo - apiStats.errorsIndex = (apiStats.errorsIndex + 1) % MAX_ERRORS - } - } -} - -func getMemoryInfo() map[string]interface{} { - var m runtime.MemStats - runtime.ReadMemStats(&m) - return map[string]interface{}{ - "total_alloc": m.TotalAlloc, - "heap_alloc": m.HeapAlloc, - "heap_inuse": m.HeapInuse, - "goroutines": runtime.NumGoroutine(), - } -} - -func healthHandler(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - - if ttsConfigErr != nil { - w.WriteHeader(http.StatusServiceUnavailable) - } else { - w.WriteHeader(http.StatusOK) - } - - apiStats.mutex.RLock() - totalRequests := apiStats.totalRequests - successfulRequests := apiStats.successfulRequests - failedRequests := apiStats.failedRequests - totalResponseTime := apiStats.totalResponseTime - recentResponseTimes := make([]float64, 0, MAX_RESPONSE_TIMES) - for _, t := range apiStats.recentResponseTimes { - if t > 0 { - recentResponseTimes = append(recentResponseTimes, t) - } - } - lastErrors := make([]string, 0, MAX_ERRORS) - for _, e := range apiStats.lastErrors { - if e != "" { - lastErrors = append(lastErrors, e) - } - } - apiStats.mutex.RUnlock() - - var errorRate float64 - if totalRequests > 0 { - errorRate = float64(failedRequests) / float64(totalRequests) * 100 - } - - var avgResponseTime float64 - if totalRequests > 0 { - avgResponseTime = totalResponseTime.Seconds() * 1000 / float64(totalRequests) - } - - envCheckStatus := checkEnvironmentVariables() - allEnvVarsSet := envCheckStatus["all_required_vars_set"].(bool) - - status := "ok" - if !allEnvVarsSet { - status = "configuration_error" - } - - response := map[string]interface{}{ - "status": status, - "service": "ByteDance TTS to OpenAI API Adapter", - "version": "2.0.0 (v3 API)", - "uptime": fmt.Sprintf("%.0f seconds", time.Since(startTime).Seconds()), - "start_time": startTime.Format(time.RFC3339), - "memory": getMemoryInfo(), - "api_stats": map[string]interface{}{ - "total_requests": totalRequests, - "successful_requests": successfulRequests, - "failed_requests": failedRequests, - "error_rate_percent": fmt.Sprintf("%.2f", errorRate), - "avg_response_time_ms": fmt.Sprintf("%.2f", avgResponseTime), - "recent_response_times_ms": recentResponseTimes, - }, - "errors": map[string]interface{}{ - "recent_errors_count": len(lastErrors), - }, - "config_status": map[string]interface{}{ - "all_required_vars_set": allEnvVarsSet, - "config_error": ttsConfigErr != nil, - }, - } - - json.NewEncoder(w).Encode(response) -} - -var startTime time.Time - -type statusRecorder struct { - http.ResponseWriter - statusCode int -} - -func (rec *statusRecorder) WriteHeader(code int) { - rec.statusCode = code - rec.ResponseWriter.WriteHeader(code) -} - -var ( - allowedOrigins []string - allowAllOrigins bool - corsMaxAgeHeader = "86400" -) - -func normalizeOrigin(origin string) string { - origin = strings.TrimSpace(origin) - origin = strings.TrimRight(origin, "/") - return strings.ToLower(origin) -} - -func initCORSConfig() { - origins := os.Getenv("ALLOWED_ORIGINS") - if origins == "" { - log.Println("警告: ALLOWED_ORIGINS 环境变量未设置") - log.Println("出于安全考虑,跨域请求将被拒绝。如需开放跨域请配置 ALLOWED_ORIGINS") - log.Println("开发环境可设置 ALLOWED_ORIGINS=* 允许所有来源(不可与凭据共用)") - return - } - - parts := strings.Split(origins, ",") - for _, p := range parts { - o := strings.TrimSpace(p) - if o == "" { - continue - } - if o == "*" { - allowAllOrigins = true - continue - } - allowedOrigins = append(allowedOrigins, normalizeOrigin(o)) - } - - if allowAllOrigins { - log.Println("警告: ALLOWED_ORIGINS=*,将允许所有来源跨域请求(不携带凭据)") - } - if len(allowedOrigins) > 0 { - log.Printf("已配置 %d 个允许的跨域来源白名单", len(allowedOrigins)) - } -} - -func checkStaticFiles() { - if _, err := os.Stat("health.html"); os.IsNotExist(err) { - log.Println("警告: health.html 不存在,/dashboard 路由将返回 404") - } -} - -func isValidOrigin(origin string) bool { - if origin == "" || origin == "null" || origin == "nil" { - return false - } - if !strings.HasPrefix(origin, "http://") && !strings.HasPrefix(origin, "https://") { - return false - } - return true -} - -func matchOrigin(origin string) (string, bool) { - if !isValidOrigin(origin) { - return "", false - } - if allowAllOrigins { - return "*", true - } - normalized := normalizeOrigin(origin) - for _, allowed := range allowedOrigins { - if allowed == normalized { - return origin, true - } - } - return "", false -} - -func corsMiddleware(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - origin := r.Header.Get("Origin") - - 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") - } - } - } - - if r.Method == http.MethodOptions { - if origin != "" { - if _, matched := matchOrigin(origin); !matched { - w.WriteHeader(http.StatusNoContent) - return - } - } - w.WriteHeader(http.StatusNoContent) - return - } - - next.ServeHTTP(w, r) - }) -} - -func main() { - startTime = time.Now() - - log.SetFlags(log.LstdFlags | log.Lshortfile) - log.SetPrefix("[TTS-Server] ") - - initAPIKeys() - initCORSConfig() - checkStaticFiles() - - ttsConfigErr = initTTSConfig() - if ttsConfigErr != nil { - log.Printf("警告: 配置初始化失败: %v", ttsConfigErr) - log.Printf("服务将继续运行,但TTS功能不可用,请检查环境变量配置") - } else { - log.Printf("配置初始化成功") - } - - router := mux.NewRouter() - - router.Use(corsMiddleware) - - router.Use(func(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) - duration := time.Since(start) - - log.Printf("%s %s %s %d %v", r.Method, r.RequestURI, r.RemoteAddr, rec.statusCode, duration) - }) - }) - - router.HandleFunc("/v1/audio/speech", openaiTTSHandler).Methods("POST", "OPTIONS") - router.HandleFunc("/health", healthHandler).Methods("GET") - router.HandleFunc("/dashboard", func(w http.ResponseWriter, r *http.Request) { - http.ServeFile(w, r, "health.html") - }).Methods("GET") - router.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { - http.Redirect(w, r, "/dashboard", http.StatusFound) - }).Methods("GET") - - port := os.Getenv("PORT") - if port == "" { - port = DEFAULT_PORT - } - - server := &http.Server{ - Addr: ":" + port, - Handler: router, - ReadTimeout: 30 * time.Second, - WriteTimeout: 120 * time.Second, - IdleTimeout: 60 * time.Second, - } - - quit := make(chan os.Signal, 1) - signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) - - go func() { - log.Printf("Starting ByteDance TTS to OpenAI API Adapter Server") - 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") - - if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed { - log.Fatalf("Server failed to start: %v", err) - } - }() - - <-quit - log.Println("Shutting down server...") - - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - if err := server.Shutdown(ctx); err != nil { - log.Printf("Server forced to shutdown: %v", err) - } else { - log.Println("Server exited gracefully") - } -}