refactor: 重构项目架构,拆分代码到模块化目录

将单文件tts_server.go重构为模块化项目结构,拆分出common、dto、middleware、router、controller、service、adapter、setting等目录,优化代码组织提升可维护性
This commit is contained in:
sun
2026-05-23 20:32:12 +08:00
parent 1b84a6c9ee
commit 4e1820d45b
21 changed files with 1536 additions and 895 deletions
+9
View File
@@ -0,0 +1,9 @@
*.exe
*.md
.env
.env.example
.git
.gitignore
tts_api_architecture.html
代码审查报告.md
fix_list.md
+4 -12
View File
@@ -9,19 +9,9 @@
BYTEDANCE_TTS_API_KEY=your_api_key_here BYTEDANCE_TTS_API_KEY=your_api_key_here
# 资源信息ID(决定使用1.0还是2.0模型) # 资源信息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 BYTEDANCE_TTS_RESOURCE_ID=seed-tts-1.0
# 发音人(音色)ID,具体参考火山引擎音色列表 # 发音人(音色)ID
# 注意:1.0音色只能搭配 seed-tts-1.0 Resource ID
# 2.0音色只能搭配 seed-tts-2.0 Resource ID
BYTEDANCE_TTS_SPEAKER=your_speaker_id_here BYTEDANCE_TTS_SPEAKER=your_speaker_id_here
# ========================================== # ==========================================
@@ -32,8 +22,10 @@ BYTEDANCE_TTS_SPEAKER=your_speaker_id_here
BYTEDANCE_TTS_TIMEOUT=30s BYTEDANCE_TTS_TIMEOUT=30s
# OpenAI兼容接口的API密钥(可选) # OpenAI兼容接口的API密钥(可选)
# 配置后,客户端请求需要携带 Authorization: Bearer <OPENAI_TTS_API_KEY>
OPENAI_TTS_API_KEY=your_openai_compatible_key_here OPENAI_TTS_API_KEY=your_openai_compatible_key_here
# CORS 跨域白名单(逗号分隔,开发环境可设 *)
# ALLOWED_ORIGINS=https://example.com,https://app.example.com
# 服务监听端口,默认8080 # 服务监听端口,默认8080
PORT=8080 PORT=8080
+26
View File
@@ -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"]
+167
View File
@@ -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
}
+19
View File
@@ -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
)
+174
View File
@@ -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
}
+21
View File
@@ -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
+31
View File
@@ -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"`
}
+40
View File
@@ -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
}
+1 -1
View File
@@ -1,4 +1,4 @@
module bytedance-tts-openai-adapter module github.com/volcano-tts/tts-api
go 1.19 go 1.19
+82
View File
@@ -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")
}
}
+59
View File
@@ -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,
},
})
}
+115
View File
@@ -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)
})
}
+35
View File
@@ -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)
})
}
+101
View File
@@ -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
}
+27
View File
@@ -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
}
+90
View File
@@ -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(),
}
}
+91
View File
@@ -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")
}
}
BIN
View File
Binary file not shown.
+444
View File
@@ -0,0 +1,444 @@
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>TTS-API 架构设计 — TTS 版 New-API</title>
<style>
*, *::before, *::after { box-sizing: border-box; margin: 0; padding: 0; }
body {
background: #fafafa; color: #111;
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", "PingFang SC", "Microsoft YaHei", monospace;
line-height: 1.75; max-width: 960px; margin: 0 auto; padding: 60px 24px 100px;
}
h1 { font-size: 2rem; font-weight: 700; letter-spacing: -0.03em; margin-bottom: 4px; }
.sub { font-size: 0.85rem; color: #888; margin-bottom: 40px; padding-bottom: 16px; border-bottom: 1px solid #ddd; }
h2 {
font-size: 1.25rem; font-weight: 700; margin: 48px 0 16px; padding-bottom: 8px;
border-bottom: 2px solid #111; letter-spacing: -0.01em;
}
h3 { font-size: 1rem; font-weight: 700; margin: 28px 0 10px; color: #333; }
p { margin-bottom: 12px; font-size: 0.9rem; color: #444; }
ul, ol { margin: 0 0 16px 20px; font-size: 0.9rem; color: #444; }
li { margin-bottom: 4px; }
/* ASCII 图表 */
.diagram {
background: #fff; border: 1px solid #ddd; padding: 20px 24px;
margin: 16px 0; overflow-x: auto; font-size: 0.78rem; line-height: 1.5;
font-family: "SF Mono", "Fira Code", "Consolas", monospace;
color: #333; white-space: pre;
}
/* 对照表 */
table {
width: 100%; border-collapse: collapse; margin: 16px 0; font-size: 0.85rem;
}
th, td {
border: 1px solid #ddd; padding: 10px 14px; text-align: left;
}
th { background: #111; color: #fff; font-weight: 600; }
tr:nth-child(even) { background: #f7f7f7; }
/* 代码块 */
.code-block {
background: #fff; border: 1px solid #ddd; padding: 16px 20px;
margin: 12px 0; font-size: 0.8rem; font-family: "SF Mono", "Fira Code", "Consolas", monospace;
overflow-x: auto; white-space: pre; color: #333;
}
/* 标签 */
.tag {
display: inline-block; background: #111; color: #fff; padding: 1px 8px;
font-size: 0.7rem; font-weight: 600; margin-right: 4px; letter-spacing: 0.03em;
}
.tag-outline { background: transparent; color: #111; border: 1px solid #111; }
.footer { margin-top: 60px; padding-top: 16px; border-top: 1px solid #ddd; font-size: 0.75rem; color: #aaa; text-align: center; }
</style>
</head>
<body>
<h1>TTS-API</h1>
<p class="sub">TTS 版 New-API 架构设计 · 多 Provider 统一 TTS 网关 · 2026-05-21</p>
<h2>1. 项目定位</h2>
<p>参考 new-api 的设计理念,TTS-API 定位为<strong>企业级 TTS 统一网关与资产管理平台</strong>,核心能力:</p>
<table>
<tr><th>能力维度</th><th>说明</th></tr>
<tr><td><strong>统一接入</strong></td><td>以 OpenAI <code>/v1/audio/speech</code> 为唯一入口,屏蔽火山引擎、阿里、Azure、讯飞等上游差异</td></tr>
<tr><td><strong>统一音色</strong></td><td>定义标准音色命名体系,自动映射到各 Provider 的实际音色 ID</td></tr>
<tr><td><strong>统一计费</strong></td><td>按字符数 / 时长计费,支持配额管理与成本核算</td></tr>
<tr><td><strong>统一治理</strong></td><td>权限分组、速率限制、渠道故障切换、审计日志、可视化看板</td></tr>
</table>
<h2>2. 与 new-api 的关键差异</h2>
<table>
<tr><th>维度</th><th>new-api(LLM)</th><th>TTS-API(本项目)</th></tr>
<tr><td>核心接口</td><td><code>/v1/chat/completions</code></td><td><code>/v1/audio/speech</code></td></tr>
<tr><td>协议转换</td><td>OpenAI ↔ Claude ↔ Gemini 互转</td><td>只需转到各 Provider 原生格式(单向)</td></tr>
<tr><td>模型映射</td><td>模型名 → 渠道选择</td><td><strong>音色映射</strong>:标准 voice → 各 Provider 实际音色 ID<br>(这是最大难点)</td></tr>
<tr><td>输出处理</td><td>文本 / JSON 流</td><td><strong>二进制音频流</strong>,需处理格式转换(wav/mp3/pcm)</td></tr>
<tr><td>计费单位</td><td>Token 数</td><td><strong>字符数 + 音频时长</strong></td></tr>
<tr><td>缓存</td><td>语义缓存(相同问题命中)</td><td><strong>音频缓存</strong>(相同文本+音色 → 直接返回已合成音频)</td></tr>
<tr><td>流式</td><td>SSE 文本流</td><td><strong>音频流式推送</strong>(边合成边返回,首字节延迟是关键指标)</td></tr>
</table>
<h2>3. 整体分层架构</h2>
<div class="diagram">
┌─────────────────────────────────────────────────────────────────┐
│ 客户端 / 应用层 │
│ 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 │
└─────────┘ └─────┘ └──────┘
</div>
<h2>4. 请求处理全流程</h2>
<div class="diagram">
客户端发送 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,流式或整段返回
</div>
<h2>5. 核心设计:适配器接口</h2>
<p>参考 new-api 的 <code>Adaptor</code> 接口设计,TTS 版适配器接口如下:</p>
<div class="code-block">
// 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 // 音色描述
}
</div>
<h2>6. 核心设计:音色映射系统</h2>
<p>这是 TTS 网关区别于 LLM 网关的<strong>最大难点与核心创新点</strong>。LLM 网关只需按模型名路由,但 TTS 需要一套跨 Provider 的音色统一体系。</p>
<div class="diagram">
标准音色命名空间
┌──────────────────────────────────┐
│ 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 │ │ │
└─────────┘ └─────────┘ └─────────┘ └─────────┘
</div>
<h3>6.1 音色映射表结构</h3>
<div class="code-block">
// 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
</div>
<h3>6.2 音色发现与自动映射建议</h3>
<p>每个 Provider 适配器实现 <code>GetVoiceList()</code>,系统启动或渠道更新时自动拉取,通过语言+性别+风格标签与标准命名空间做相似度匹配,自动生成映射建议,管理员在 Web UI 审核确认即可。</p>
<h2>7. 渠道调度与故障切换</h2>
<div class="diagram">
请求 → 查找"tts-1-gentle-male"可用的渠道列表
│
▼
按 priority 排序 → 加权随机选择一个渠道
│
▼
适配器发送请求 → 成功?
├── 是 → 返回音频,记录成功
│
└── 否 → 标记渠道失败
│
▼
自动切换 priority+1 的渠道重试
│
▼
所有渠道失败 → 返回 503 + 错误详情
</div>
<h2>8. 项目目录结构</h2>
<div class="code-block">
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
</div>
<h2>9. 管理后台页面规划</h2>
<table>
<tr><th>页面</th><th>功能</th></tr>
<tr><td><strong>Dashboard</strong></td><td>今日合成次数、字符数、时长、费用、渠道健康状态、QPS 曲线</td></tr>
<tr><td><strong>渠道管理</strong></td><td>添加/编辑 Provider(API Key、Resource ID、权重、并发上限)</td></tr>
<tr><td><strong>音色映射</strong></td><td>核心页面:标准音色 ↔ 各渠道音色 ID 的映射表,支持自动匹配建议与手动调整</td></tr>
<tr><td><strong>API Key</strong></td><td>生成/管理用户 API Key,绑定分组与配额</td></tr>
<tr><td><strong>用户管理</strong></td><td>用户 CRUD、分组、角色(Admin / User)</td></tr>
<tr><td><strong>计费统计</strong></td><td>按用户/渠道/日期维度的用量与费用报表</td></tr>
<tr><td><strong>调用日志</strong></td><td>每次 TTS 请求的详细日志(文本、音色、渠道、耗时、费用)</td></tr>
</table>
<h2>10. 一期 vs 二期路线图</h2>
<h3>一期(MVP,你现有的 Volcano-Engine-TTS-UI 升级版)</h3>
<table>
<tr><th>模块</th><th>内容</th></tr>
<tr><td>适配器</td><td>火山引擎 + Azure + 阿里云,3 个 Provider</td></tr>
<tr><td>接口</td><td><code>/v1/audio/speech</code>,OpenAI 兼容</td></tr>
<tr><td>音色映射</td><td>硬编码映射表(配置文件),先跑通再抽象</td></tr>
<tr><td>计费</td><td>简单字符数计数 + 日志</td></tr>
<tr><td>管理后台</td><td>极简版:渠道配置页面 + 用量看板</td></tr>
<tr><td>数据库</td><td>SQLite,单文件部署</td></tr>
</table>
<h3>二期(完整版)</h3>
<table>
<tr><th>模块</th><th>内容</th></tr>
<tr><td>适配器</td><td>扩到 8+ Provider(讯飞、腾讯、Fish Audio、Bert-VITS2、OpenAI)</td></tr>
<tr><td>音色映射</td><td>数据库驱动 + Web UI 可视化管理 + 自动匹配建议</td></tr>
<tr><td>计费</td><td>字符数/时长双维度计费,用户配额,欠费阻断</td></tr>
<tr><td>缓存</td><td>Redis 音频缓存,相同文本+音色直接命中</td></tr>
<tr><td>流式</td><td>支持流式 TTS(SSE 推送音频 chunk),降低首字节延迟</td></tr>
<tr><td>管理后台</td><td>完整 React 后台(参考 new-api 的 Semi Design UI)</td></tr>
<tr><td>数据库</td><td>MySQL / PostgreSQL 支持</td></tr>
<tr><td>部署</td><td>Docker Compose 一键部署</td></tr>
</table>
<h2>11. 关键工程建议</h2>
<ol>
<li><strong>从你现有的火山引擎适配器起步</strong>,先抽象出 <code>TTSAdapter</code> 接口,接入 2~3 个 Provider 验证接口设计是否合理,不要一上来就搞 8 个适配器。</li>
<li><strong>音色映射先硬编码</strong>,跑通流程后再做成数据库驱动的 Web UI。映射表是长期维护工作,需要社区共建。</li>
<li><strong>音频格式转换用 ffmpeg</strong>,Go 侧通过 <code>exec.Command</code> 调用或使用 <code>go-ffmpeg</code> 绑定。各 Provider 输出格式不同(wav / mp3 / pcm),统一转码是刚需。</li>
<li><strong>渠道调度直接复用 new-api 的加权随机 + 故障重试思路</strong>,这是成熟的模式,不需要重新发明。</li>
<li><strong>管理后台前期可以不写</strong>,SQLite + 配置文件就能用;等 Provider 多了再补 React 前端。</li>
<li><strong>考虑直接 fork new-api 改造</strong>:new-api 的渠道管理、用户系统、计费框架、中间件、部署方案都是现成的,你只需要把 relay 层的 LLM 适配器替换成 TTS 适配器,再加音色映射模块。这比从零搭建快得多。</li>
</ol>
<div class="footer">Architecture Design for TTS-API · Inspired by new-api (QuantumNous)</div>
</body>
</html>
-882
View File
@@ -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")
}
}