Compare commits
31
Commits
v0.1.0
...
7e2d050d51
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7e2d050d51 | ||
|
|
f39c72acbe | ||
|
|
7e1102902d | ||
|
|
746da76fa4 | ||
|
|
3dc9632c1b | ||
|
|
0aad65ed78 | ||
|
|
21b86bfcfe | ||
|
|
82cc68e7ee | ||
|
|
361a9d6401 | ||
|
|
15b0470cc8 | ||
|
|
bcbd796fa5 | ||
|
|
8592843bdf | ||
|
|
b93ede29e0 | ||
|
|
4aed9667b7 | ||
|
|
03bb98beb8 | ||
|
|
bc42295ff6 | ||
|
|
f704e7d71d | ||
|
|
61431e00ba | ||
|
|
3d50b6c69d | ||
|
|
977e9ccadb | ||
|
|
b92d3dbc00 | ||
|
|
b92973cdc9 | ||
|
|
be6c2ad34e | ||
|
|
23b962a90e | ||
|
|
9b2a1d1531 | ||
|
|
9c35f780db | ||
|
|
7bedb222d1 | ||
|
|
4e1820d45b | ||
|
|
1b84a6c9ee | ||
|
|
4c93638250 | ||
|
|
45591a4e3a |
@@ -0,0 +1,9 @@
|
||||
*.exe
|
||||
*.md
|
||||
.env
|
||||
.env.example
|
||||
.git
|
||||
.gitignore
|
||||
tts_api_architecture.html
|
||||
代码审查报告.md
|
||||
fix_list.md
|
||||
+11
-12
@@ -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
|
||||
|
||||
# ==========================================
|
||||
@@ -31,9 +21,18 @@ BYTEDANCE_TTS_SPEAKER=your_speaker_id_here
|
||||
# 请求超时时间,默认30秒
|
||||
BYTEDANCE_TTS_TIMEOUT=30s
|
||||
|
||||
# 音频格式:mp3/ogg_opus/pcm/wav(默认mp3)
|
||||
# 注意:流式场景下wav会多次返回header,内部自动用pcm请求再封装header
|
||||
BYTEDANCE_TTS_FORMAT=mp3
|
||||
|
||||
# 音频采样率:8000/16000/22050/24000/32000/44100/48000(默认24000)
|
||||
BYTEDANCE_TTS_SAMPLE_RATE=24000
|
||||
|
||||
# OpenAI兼容接口的API密钥(可选)
|
||||
# 配置后,客户端请求需要携带 Authorization: Bearer <OPENAI_TTS_API_KEY>
|
||||
OPENAI_TTS_API_KEY=your_openai_compatible_key_here
|
||||
|
||||
# CORS 跨域白名单(逗号分隔,开发环境可设 *)
|
||||
# ALLOWED_ORIGINS=https://example.com,https://app.example.com
|
||||
|
||||
# 服务监听端口,默认8080
|
||||
PORT=8080
|
||||
|
||||
+14
@@ -0,0 +1,14 @@
|
||||
# Go build cache
|
||||
.gocache/
|
||||
*.exe
|
||||
*.test
|
||||
*.out
|
||||
|
||||
# Editor / OS
|
||||
.vscode/
|
||||
.idea/
|
||||
.DS_Store
|
||||
Thumbs.db
|
||||
|
||||
# Logs
|
||||
*.log
|
||||
+31
@@ -0,0 +1,31 @@
|
||||
FROM golang:1.26-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.21
|
||||
|
||||
RUN apk --no-cache add ca-certificates tzdata \
|
||||
&& addgroup -S appgroup && adduser -S appuser -G appgroup
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY --from=builder /app/tts-api .
|
||||
COPY --from=builder /app/health.html .
|
||||
|
||||
RUN chown -R appuser:appgroup /app
|
||||
|
||||
USER appuser
|
||||
|
||||
EXPOSE 8080
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=5s --start-period=5s --retries=3 \
|
||||
CMD wget -qO- http://localhost:8080/health || exit 1
|
||||
|
||||
ENTRYPOINT ["./tts-api"]
|
||||
@@ -6,20 +6,15 @@
|
||||
|
||||
### 主要特性
|
||||
|
||||
- ✅ 完全兼容OpenAI `/v1/audio/speech` API接口
|
||||
- ✅ 支持火山引擎TTS v3 API(单向流式)
|
||||
- ✅ 支持API Key鉴权方式
|
||||
- ✅ 支持多种发音人和模型版本
|
||||
- ✅ 内置速率限制和统计功能
|
||||
- ✅ 支持配置API密钥验证
|
||||
- ✅ 并发限制:最多同时处理10个请求(保护上游API)
|
||||
- ✅ 跨平台支持(Windows/Linux/macOS)
|
||||
|
||||
## 文件说明
|
||||
|
||||
- `tts_server.go` - 主程序源码
|
||||
- `.env.example` - 环境变量配置示例
|
||||
- `go.mod` / `go.sum` - Go模块依赖
|
||||
- 完全兼容OpenAI `/v1/audio/speech` API接口
|
||||
- 支持火山引擎TTS v3 API(HTTP Chunked单向流式)
|
||||
- 支持多种音频格式:mp3、ogg_opus、pcm、wav
|
||||
- 支持API Key鉴权方式
|
||||
- 支持多种发音人和模型版本
|
||||
- 内置速率限制和统计功能
|
||||
- 支持配置API密钥验证
|
||||
- 并发限制:最多同时处理10个请求(保护上游API)
|
||||
- 跨平台支持(Windows/Linux/macOS)
|
||||
|
||||
## 快速开始
|
||||
|
||||
@@ -31,7 +26,7 @@
|
||||
### 1. 编译程序
|
||||
|
||||
```bash
|
||||
go build -o tts_server tts_server.go
|
||||
go build -o tts-api .
|
||||
```
|
||||
|
||||
### 2. 配置环境变量
|
||||
@@ -48,10 +43,10 @@ cp .env.example .env
|
||||
|
||||
```bash
|
||||
# Windows
|
||||
tts_server.exe
|
||||
tts-api.exe
|
||||
|
||||
# Linux/macOS
|
||||
./tts_server
|
||||
./tts-api
|
||||
```
|
||||
|
||||
服务默认监听 `8080` 端口。
|
||||
@@ -70,9 +65,13 @@ tts_server.exe
|
||||
|
||||
| 变量名 | 说明 | 默认值 |
|
||||
|--------|------|--------|
|
||||
| `BYTEDANCE_TTS_MODEL` | 模型子版本(复刻音色必填) | `seed-tts-2.0-standard` |
|
||||
| `BYTEDANCE_TTS_TIMEOUT` | 请求超时时间 | `30s` |
|
||||
| `BYTEDANCE_TTS_FORMAT` | 音频格式:`mp3` / `ogg_opus` / `pcm` / `wav` | `mp3` |
|
||||
| `BYTEDANCE_TTS_SAMPLE_RATE` | 采样率:8000/16000/22050/24000/32000/44100/48000 | `24000` |
|
||||
| `OPENAI_TTS_API_KEY` | OpenAI兼容接口的API密钥(逗号分隔支持多个) | 无 |
|
||||
| `PORT` | 服务监听端口 | `8080` |
|
||||
| `ALLOWED_ORIGINS` | 允许跨域请求的来源(多个用英文逗号分隔;调试可设为 `*`) | 无(不设则拒绝所有跨域) |
|
||||
|
||||
### Resource ID 说明
|
||||
|
||||
@@ -85,8 +84,53 @@ tts_server.exe
|
||||
| `seed-icl-1.0-concurr` | 声音复刻1.0并发版 |
|
||||
| `seed-icl-2.0` | 声音复刻2.0字符版 |
|
||||
|
||||
> 上表为通用模型名。火山控制台实际显示的资源 ID 字符串格式通常是 `volc.megatts.default`、`volc.megatts.icl` 等(带版本号会形如 `volc.megatts.icl.2_0`),**以控制台资源管理页面显示的字符串为准**。资源 ID 与音色必须**同时在控制台开通**才能组合使用,否则 API 会返回 `code=55000000, message=resource ID is mismatched with speaker related resource`。
|
||||
|
||||
**注意:** 1.0音色只能搭配 `seed-tts-1.0` Resource ID,2.0音色只能搭配 `seed-tts-2.0` Resource ID。
|
||||
|
||||
### 音频格式说明
|
||||
|
||||
本项目使用火山 v3 **HTTP Chunked 单向流式** API([官方文档](https://www.volcengine.com/docs/6561/1598757)),支持以下音频格式:
|
||||
|
||||
| 格式 | Content-Type | 说明 |
|
||||
|------|-------------|------|
|
||||
| `mp3` | `audio/mpeg` | 默认格式,API原生支持,推荐使用 |
|
||||
| `ogg_opus` | `audio/ogg` | OGG Opus格式,API原生支持 |
|
||||
| `pcm` | `audio/pcm` | 原始PCM数据,API原生支持 |
|
||||
| `wav` | `audio/wav` | 本端用PCM请求API后封装WAV header(流式场景下API的wav会多次返回header,所以内部用pcm再拼装) |
|
||||
|
||||
> 流式场景下直接请求 wav 格式,API每个chunk都会返回一个完整的 wav header,导致拼接后的音频损坏。因此当用户选择 wav 输出时,适配器自动用 pcm 格式请求 API,最后在客户端拼装标准 44 字节 WAV 文件头。
|
||||
|
||||
### v3 API 调用说明
|
||||
|
||||
本项目按火山 v3 HTTP Chunked 单向流式 TTS API 实现([官方文档](https://www.volcengine.com/docs/6561/1598757)),相比 v1/v2 有以下关键差异:
|
||||
|
||||
- **协议**:HTTP Chunked 流式,请求路径 `https://openspeech.bytedance.com/api/v3/tts/unidirectional`
|
||||
- **不再使用业务集群**(`cluster` 字段在 v3 已废弃),改用 `X-Api-Resource-Id` HTTP header 路由模型
|
||||
- **鉴权 header 只有** `X-Api-Key` 一个,无 `Authorization`,无 app 对象
|
||||
- **用量返回**:携带 `X-Control-Require-Usage-Tokens-Return: *` header,合成结束时响应中包含 `usage` 字段
|
||||
- **`req_params.model` 字段**:v3 必须显式传子模型版本。可选值:
|
||||
- `seed-tts-2.0-standard`(默认,标准版,常规音色/复刻音色通用)
|
||||
- `seed-tts-2.0-expressive`(表现力增强版,部分复刻音色推荐)
|
||||
- 留空时会用 `seed-tts-2.0-standard` 作为兜底
|
||||
- **复刻音色(`S_` 开头的 speaker)必须显式传 model**,否则可能因默认模型与复刻音色不匹配返回 `55000000`
|
||||
- **响应 event 字段**:`TTSSentenceStart`/`TTSSentenceEnd` 标记句子边界,音频数据在默认 event 中返回
|
||||
|
||||
## CORS 跨域配置
|
||||
|
||||
跨域请求由 `ALLOWED_ORIGINS` 环境变量控制,按**完整 origin**(含协议 + 域名 + 端口)精确匹配:
|
||||
|
||||
- `https://app.example.com` — 精确匹配一个来源
|
||||
- `https://a.com,https://b.com` — 多个来源英文逗号分隔
|
||||
- `*` — 允许所有来源(**不可与凭据请求共存**,需同时去掉 `Authorization` 头)
|
||||
- `app.example.com` — 缺协议头,**永远不会匹配**(服务端强制校验 `http://` / `https://` 开头)
|
||||
|
||||
**典型坑**:
|
||||
|
||||
1. 客户端 URL 是 `http://` 但服务端是 `https://`:浏览器按 `http://...` 的 origin 发请求,白名单里的 `https://...` 不会匹配 → 403。**客户端必须用 `https://` 开头**。
|
||||
2. `ALLOWED_ORIGINS=*` + 客户端带 `Authorization`:浏览器按规范会**直接拒绝预检**(凭据 + 通配符冲突),POST 根本发不出去。
|
||||
3. 同源请求(前端和 TTS 服务同域名)不受 CORS 限制,`ALLOWED_ORIGINS` 怎么配都不影响。
|
||||
|
||||
## API 使用说明
|
||||
|
||||
### OpenAI 兼容接口
|
||||
@@ -103,7 +147,7 @@ tts_server.exe
|
||||
"model": "tts-1",
|
||||
"input": "你好,这是一个测试文本",
|
||||
"voice": "alloy",
|
||||
"response_format": "wav",
|
||||
"response_format": "mp3",
|
||||
"speed": 1.0
|
||||
}
|
||||
```
|
||||
@@ -111,16 +155,33 @@ tts_server.exe
|
||||
**参数说明:**
|
||||
- `model` - 模型名称(OpenAI兼容,实际不影响)
|
||||
- `input` - 要合成的文本
|
||||
- `voice` - 发音人(OpenAI兼容,实际不影响)
|
||||
- `response_format` - 输出格式:仅支持 `wav`
|
||||
- `voice` - 发音人(OpenAI兼容,实际使用配置的BYTEDANCE_TTS_SPEAKER)
|
||||
- `response_format` - 输出格式:`mp3`(默认)、`opus`(映射到ogg_opus)、`wav`、`pcm`、`aac`/`flac`(降级到mp3)
|
||||
- `speed` - 语速:0.25 ~ 4.0
|
||||
|
||||
**格式映射(OpenAI → 火山):**
|
||||
|
||||
| OpenAI response_format | 火山 API 格式 | Content-Type |
|
||||
|------------------------|--------------|--------------|
|
||||
| `mp3` | mp3 | audio/mpeg |
|
||||
| `opus` | ogg_opus | audio/ogg |
|
||||
| `wav` | pcm → 封装wav header | audio/wav |
|
||||
| `pcm` | pcm | audio/pcm |
|
||||
| `aac` / `flac` | mp3(降级) | audio/mpeg |
|
||||
|
||||
**示例调用:**
|
||||
|
||||
```bash
|
||||
# MP3格式(默认)
|
||||
curl -X POST "http://localhost:8080/v1/audio/speech" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"model":"tts-1","input":"你好,世界","voice":"alloy","speed":1.0}' \
|
||||
-o output.mp3
|
||||
|
||||
# WAV格式
|
||||
curl -X POST "http://localhost:8080/v1/audio/speech" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"model":"tts-1","input":"你好,世界","voice":"alloy","response_format":"wav","speed":1.0}' \
|
||||
-o output.wav
|
||||
```
|
||||
|
||||
@@ -184,10 +245,10 @@ curl http://localhost:8080/health
|
||||
|
||||
```bash
|
||||
# Windows
|
||||
set PORT=8081 && tts_server.exe
|
||||
set PORT=8081 && tts-api.exe
|
||||
|
||||
# Linux/macOS
|
||||
PORT=8081 ./tts_server
|
||||
PORT=8081 ./tts-api
|
||||
```
|
||||
|
||||
### 3. 如何配置多个API密钥?
|
||||
@@ -198,13 +259,72 @@ PORT=8081 ./tts_server
|
||||
OPENAI_TTS_API_KEY=sk-key1,sk-key2,sk-key3
|
||||
```
|
||||
|
||||
### 4. 查看日志
|
||||
### 4. 出现 `code=55000000, message=resource ID is mismatched with speaker related resource` 怎么办?
|
||||
|
||||
服务启动后会输出详细日志,包括:
|
||||
- 服务启动信息
|
||||
- 配置状态
|
||||
- 请求统计信息
|
||||
- 错误详情
|
||||
这是火山引擎 v3 API 返回的**资源/音色不匹配**错误,不是网络或超时问题。修复方法:
|
||||
|
||||
1. 去**火山控制台** → 语音技术 → 你的应用 → 资源管理或音色库
|
||||
2. 用控制台的在线体验/调试试一下同一对 `BYTEDANCE_TTS_RESOURCE_ID` + 音色
|
||||
3. 控制台能合成的组合才是正确的
|
||||
4. 把控制台显示的**实际资源 ID 字符串**(通常是 `volc.megatts.*` 格式)填到 `BYTEDANCE_TTS_RESOURCE_ID`
|
||||
5. 如果你用的是**声音复刻**音色(speaker 以 `S_` 开头),同时确认设置了 `BYTEDANCE_TTS_MODEL`(推荐 `seed-tts-2.0-standard` 或 `seed-tts-2.0-expressive`)。复刻音色不传 `model` 字段是 55000000 的常见原因之一
|
||||
|
||||
### 5. PowerShell 下 `curl` 命令被解释错
|
||||
|
||||
PowerShell 里 `curl` 是 `Invoke-WebRequest` 的别名,参数完全不同。**必须写 `curl.exe`**:
|
||||
|
||||
```powershell
|
||||
curl.exe -v -X POST "http://localhost:8080/v1/audio/speech" -H "Content-Type: application/json" -H "Authorization: Bearer YOUR_KEY" --data-binary "@body.json"
|
||||
```
|
||||
|
||||
另外 PowerShell 里 `{"foo":"bar"}` 不加单引号会被当成脚本块解析。**要么用单引号包 JSON**,要么把 body 写到文件用 `--data-binary "@file.json"`。
|
||||
|
||||
### 6. WAV 格式音频播放异常?
|
||||
|
||||
流式场景下火山 API 的 wav 格式会每个 chunk 都返回完整的 wav header,拼接后音频损坏。本项目已自动处理:选择 wav 输出时,内部用 pcm 格式请求 API,最后拼装标准 wav header。如仍有问题,建议改用 `mp3` 格式。
|
||||
|
||||
### 7. 查看日志
|
||||
|
||||
服务启动后输出到 stdout/stderr。常见日志关键字:
|
||||
|
||||
**中间件层拒绝**(有专门日志):
|
||||
|
||||
```
|
||||
CORS拦截: 来源="https://..." 路径=/v1/audio/speech 方法=POST 客户端=...
|
||||
警告: 已超过IP速率限制,拒绝请求 - 客户端IP: 1.2.3.4
|
||||
警告: 已达到最大并发请求数限制,拒绝请求 - 客户端IP: 1.2.3.4
|
||||
```
|
||||
|
||||
**Controller 层拒绝**(每条都带具体原因和客户端 IP):
|
||||
|
||||
```
|
||||
警告: 错误的方法 - 方法=GET 期望=POST 路径=/v1/audio/speech 客户端=...
|
||||
警告: API Key 鉴权失败 - 路径=/v1/audio/speech 客户端=... 远端=...
|
||||
警告: TTS配置未就绪,拒绝请求 - 错误=缺少必需的环境变量: [BYTEDANCE_TTS_API_KEY] 路径=...
|
||||
警告: 请求体过大 - 路径=... 限制=1048576字节
|
||||
警告: 读取请求体失败 - 路径=... 错误=...
|
||||
警告: JSON 解析失败 - 路径=... 错误=... body前200字节="..."
|
||||
警告: Model 名过长 - 路径=... 长度=80 限制=64
|
||||
警告: Model 名含非法字符 - 路径=... model前50字节="..."
|
||||
警告: input 字段为空 - 路径=...
|
||||
警告: input 文本过长 - 路径=... 长度=6000 限制=5000
|
||||
警告: TTS 合成失败 - 路径=... 文本长度=50 耗时=114ms 错误=...
|
||||
```
|
||||
|
||||
**适配器层日志**:
|
||||
|
||||
```
|
||||
Sentence start: sequence=0, sentence=...
|
||||
Sentence end: sequence=0
|
||||
TTS synthesis completed, usage: &{TextWords:5}
|
||||
```
|
||||
|
||||
**请求结束通用日志**(每个请求都有,由 Logger 中间件输出):
|
||||
|
||||
```
|
||||
POST /v1/audio/speech 1.2.3.4:56789 200 245ms
|
||||
POST /v1/audio/speech 1.2.3.4:56789 400 1ms
|
||||
```
|
||||
|
||||
## 部署建议
|
||||
|
||||
@@ -222,7 +342,7 @@ Type=simple
|
||||
User=www-data
|
||||
WorkingDirectory=/www/wwwroot/tts-server
|
||||
EnvironmentFile=/www/wwwroot/tts-server/.env
|
||||
ExecStart=/www/wwwroot/tts-server/tts_server
|
||||
ExecStart=/www/wwwroot/tts-server/tts-api
|
||||
Restart=always
|
||||
RestartSec=10
|
||||
|
||||
@@ -238,6 +358,14 @@ sudo systemctl enable tts-server
|
||||
sudo systemctl start tts-server
|
||||
```
|
||||
|
||||
### Docker 部署
|
||||
|
||||
```bash
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
环境变量通过 `.env` 文件或 docker-compose.yml 传入。
|
||||
|
||||
## 许可证
|
||||
|
||||
本项目采用非商业用途许可协议。您可以免费使用本软件用于非商业目的,但禁止用于任何商业活动。详细条款请参阅 [LICENSE](LICENSE) 文件。
|
||||
@@ -249,3 +377,8 @@ sudo systemctl start tts-server
|
||||
2. 网络是否能访问火山引擎TTS服务
|
||||
3. 鉴权信息是否有效
|
||||
4. Resource ID与Speaker是否匹配
|
||||
5. ALLOWED_ORIGINS 是否包含前端完整 origin(含 https://)
|
||||
6. 客户端请求 URL 是否以 https:// 开头
|
||||
7. 生产环境凭据是否定期轮换(API Key 明文出现在日志/对话中时立刻重置)
|
||||
8. 复刻音色(speaker 以 `S_` 开头)是否设置了 `BYTEDANCE_TTS_MODEL`(默认 `seed-tts-2.0-standard`)
|
||||
9. 音频格式是否匹配客户端解码能力(默认 mp3 兼容性最好)
|
||||
|
||||
@@ -0,0 +1,322 @@
|
||||
package volcano
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/volcano-tts/tts-api/dto"
|
||||
)
|
||||
|
||||
// 火山引擎 TTS v3 HTTP 单向流式 API 客户端。
|
||||
// 官方文档:https://www.volcengine.com/docs/6561/1598757
|
||||
// 官方 Go 示例:请求体仅含 req_params;本实现按 commit 4aed966 经验额外带上
|
||||
// user.uid 和 namespace="UnidirectionalTTS"(早期用其它 namespace 出现过兼容性
|
||||
// 问题,显式指定最稳)。复刻音色场景额外带 req_params.model。
|
||||
|
||||
type HTTPClient struct {
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
func NewHTTPClient() *HTTPClient {
|
||||
return &HTTPClient{
|
||||
client: &http.Client{
|
||||
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)
|
||||
}
|
||||
|
||||
// --- 请求体结构(对应火山 v3 API 请求 JSON) ---
|
||||
|
||||
type ttsRequest struct {
|
||||
User ttsUser `json:"user"`
|
||||
Namespace string `json:"namespace"`
|
||||
ReqParams ttsReqParams `json:"req_params"`
|
||||
}
|
||||
|
||||
type ttsUser struct {
|
||||
UID string `json:"uid"`
|
||||
}
|
||||
|
||||
type ttsReqParams struct {
|
||||
Text string `json:"text"`
|
||||
Speaker string `json:"speaker"`
|
||||
Model string `json:"model"`
|
||||
AudioParams ttsAudioParams `json:"audio_params"`
|
||||
}
|
||||
|
||||
type ttsAudioParams struct {
|
||||
Format string `json:"format"`
|
||||
SampleRate int `json:"sample_rate"`
|
||||
SpeechRate int `json:"speech_rate"`
|
||||
}
|
||||
|
||||
// convertSpeedToSpeechRate 把 OpenAI 风格的 speed(倍率)转成火山 v3 的 speech_rate(百分比)。
|
||||
// 文档规定 speech_rate 范围 [-50, 100],对应 0.5x ~ 2.0x 倍速。
|
||||
// 输入超出范围会被截断到边界值。
|
||||
func convertSpeedToSpeechRate(speed float64) int {
|
||||
rate := int((speed - 1.0) * 100)
|
||||
if rate < -50 {
|
||||
rate = -50
|
||||
}
|
||||
if rate > 100 {
|
||||
rate = 100
|
||||
}
|
||||
return rate
|
||||
}
|
||||
|
||||
// resolveAPIFormat 根据用户期望的输出格式决定实际请求火山 API 的格式。
|
||||
// 文档明确指出:流式场景下传入 wav 会多次返回 wav header,建议使用 pcm。
|
||||
// 因此当用户要 wav 输出时,用 pcm 请求 API,最后由本端拼装完整 wav header。
|
||||
func resolveAPIFormat(desiredFormat string) (apiFormat string, needWavHeader bool) {
|
||||
switch desiredFormat {
|
||||
case "wav":
|
||||
return "pcm", true
|
||||
case "mp3", "ogg_opus", "pcm":
|
||||
return desiredFormat, false
|
||||
default:
|
||||
return "mp3", false
|
||||
}
|
||||
}
|
||||
|
||||
// buildWavHeader 构造标准 44 字节 WAV 文件头(16-bit PCM, mono)。
|
||||
func buildWavHeader(dataLen int, sampleRate int) []byte {
|
||||
header := make([]byte, 44)
|
||||
byteRate := sampleRate * 2 // 16bit * 1channel / 8 * sampleRate
|
||||
blockAlign := 2 // 16bit / 8 * 1channel
|
||||
|
||||
copy(header[0:4], "RIFF")
|
||||
binary.LittleEndian.PutUint32(header[4:8], uint32(36+dataLen))
|
||||
copy(header[8:12], "WAVE")
|
||||
copy(header[12:16], "fmt ")
|
||||
binary.LittleEndian.PutUint32(header[16:20], 16) // SubChunk1Size
|
||||
binary.LittleEndian.PutUint16(header[20:22], 1) // PCM format
|
||||
binary.LittleEndian.PutUint16(header[22:24], 1) // NumChannels
|
||||
binary.LittleEndian.PutUint32(header[24:28], uint32(sampleRate))
|
||||
binary.LittleEndian.PutUint32(header[28:32], uint32(byteRate))
|
||||
binary.LittleEndian.PutUint16(header[32:34], uint16(blockAlign))
|
||||
binary.LittleEndian.PutUint16(header[34:36], 16) // BitsPerSample
|
||||
copy(header[36:40], "data")
|
||||
binary.LittleEndian.PutUint32(header[40:44], uint32(dataLen))
|
||||
|
||||
return header
|
||||
}
|
||||
|
||||
// FormatContentType 返回音频格式对应的 HTTP Content-Type。
|
||||
func FormatContentType(format string) string {
|
||||
switch format {
|
||||
case "mp3":
|
||||
return "audio/mpeg"
|
||||
case "wav":
|
||||
return "audio/wav"
|
||||
case "ogg_opus":
|
||||
return "audio/ogg"
|
||||
case "pcm":
|
||||
return "audio/pcm"
|
||||
default:
|
||||
return "application/octet-stream"
|
||||
}
|
||||
}
|
||||
|
||||
// MapOpenAIFormat 将 OpenAI TTS response_format 映射为火山 API 支持的格式。
|
||||
// OpenAI 支持: mp3, opus, aac, flac, wav, pcm
|
||||
// 火山支持: mp3, ogg_opus, pcm, wav(流式不推荐)
|
||||
func MapOpenAIFormat(openaiFormat string) string {
|
||||
switch openaiFormat {
|
||||
case "mp3":
|
||||
return "mp3"
|
||||
case "opus":
|
||||
return "ogg_opus"
|
||||
case "wav":
|
||||
return "wav"
|
||||
case "pcm":
|
||||
return "pcm"
|
||||
case "aac", "flac":
|
||||
return "mp3" // 火山不支持 aac/flac,降级到 mp3
|
||||
default:
|
||||
return "mp3"
|
||||
}
|
||||
}
|
||||
|
||||
func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text string, speed float64, voice string, requestFormat string) (*dto.SynthesisResult, error) {
|
||||
reqID := uuid.NewString()
|
||||
speechRate := convertSpeedToSpeechRate(speed)
|
||||
|
||||
speaker := config.Speaker
|
||||
if voice != "" {
|
||||
speaker = voice
|
||||
}
|
||||
|
||||
model := config.Model
|
||||
if model == "" {
|
||||
model = "seed-tts-2.0-standard" // 文档默认值,复刻音色可设为 seed-tts-2.0-expressive
|
||||
}
|
||||
|
||||
// 决定实际输出格式:优先用请求中指定的格式,否则用配置中的格式,最后默认 mp3
|
||||
outputFormat := config.Format
|
||||
if requestFormat != "" {
|
||||
outputFormat = requestFormat
|
||||
}
|
||||
if outputFormat == "" {
|
||||
outputFormat = "mp3"
|
||||
}
|
||||
|
||||
// 根据输出格式确定 API 请求格式(wav → pcm + 本端封装 header)
|
||||
apiFormat, needWavHeader := resolveAPIFormat(outputFormat)
|
||||
|
||||
sampleRate := config.SampleRate
|
||||
if sampleRate == 0 {
|
||||
sampleRate = 24000
|
||||
}
|
||||
|
||||
// 构造请求体:严格按 v3 API 文档 JSON 结构
|
||||
req := ttsRequest{
|
||||
User: ttsUser{UID: reqID},
|
||||
Namespace: "UnidirectionalTTS",
|
||||
ReqParams: ttsReqParams{
|
||||
Text: text,
|
||||
Speaker: speaker,
|
||||
Model: model,
|
||||
AudioParams: ttsAudioParams{
|
||||
Format: apiFormat,
|
||||
SampleRate: sampleRate,
|
||||
SpeechRate: speechRate,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
body, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal TTS request: %w", err)
|
||||
}
|
||||
// 诊断日志:记录实际发到上游的请求体(去 model/speaker/resource 关键字段)
|
||||
log.Printf("TTS upstream request: X-Api-Resource-Id=%s speaker=%s model=%s namespace=UnidirectionalTTS body=%s",
|
||||
config.ResourceId, speaker, model, string(body))
|
||||
|
||||
// 鉴权 header 按 v3 新版控制台方式(Connection 由 Go http 默认 keep-alive)
|
||||
headers := map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
// 请求用量返回,合成结束时响应中携带 usage 字段
|
||||
"X-Control-Require-Usage-Tokens-Return": "*",
|
||||
"X-Api-Resource-Id": config.ResourceId, // 模型路由(seed-tts-2.0 / seed-icl-2.0)
|
||||
"X-Api-Request-Id": reqID,
|
||||
"X-Api-Key": config.ApiKey, // v3 鉴权 key
|
||||
}
|
||||
|
||||
resp, err := httpClient.PostStream(config.URL, headers, body, config.Timeout)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("send TTS request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
respBody, readErr := io.ReadAll(resp.Body)
|
||||
if readErr != nil {
|
||||
log.Printf("TTS service error: status=%d, read body fail: %v", resp.StatusCode, readErr)
|
||||
} else {
|
||||
log.Printf("TTS service error: status=%d, body=%s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
return nil, fmt.Errorf("TTS service error: status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
var audioData []byte
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
// 初始 1MB / 最大 8MB,与示例同量级,留足 TTS 长文本 room
|
||||
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
|
||||
}
|
||||
|
||||
// code=20000000 表示合成结束
|
||||
if v3Resp.Code == 20000000 {
|
||||
if v3Resp.Usage != nil {
|
||||
log.Printf("TTS synthesis completed, usage: %+v", v3Resp.Usage)
|
||||
}
|
||||
// 跳过后续可能的空行
|
||||
for scanner.Scan() {
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
// 非零 code 为错误
|
||||
if v3Resp.Code != 0 {
|
||||
log.Printf("TTS service error: code=%d, message=%s, event=%s", v3Resp.Code, v3Resp.Message, v3Resp.Event)
|
||||
return nil, fmt.Errorf("TTS service error: %s", v3Resp.Message)
|
||||
}
|
||||
|
||||
// 根据 event 字段分类处理
|
||||
switch v3Resp.Event {
|
||||
case "TTSSentenceStart":
|
||||
log.Printf("Sentence start: sequence=%d, sentence=%s", v3Resp.Sequence, v3Resp.Sentence)
|
||||
case "TTSSentenceEnd":
|
||||
log.Printf("Sentence end: sequence=%d", v3Resp.Sequence)
|
||||
default:
|
||||
// 音频数据 chunk:data 字段为 base64 编码的音频片段
|
||||
if v3Resp.Data != "" {
|
||||
chunk, err := base64.StdEncoding.DecodeString(v3Resp.Data)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode audio chunk: %w", err)
|
||||
}
|
||||
audioData = append(audioData, chunk...)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if err := scanner.Err(); err != nil {
|
||||
return nil, fmt.Errorf("read TTS stream: %w", err)
|
||||
}
|
||||
|
||||
if len(audioData) == 0 {
|
||||
return nil, fmt.Errorf("no audio data received from TTS service")
|
||||
}
|
||||
|
||||
// 若输出格式为 wav,需要在 pcm 数据前拼装完整的 wav header
|
||||
if needWavHeader {
|
||||
wavHeader := buildWavHeader(len(audioData), sampleRate)
|
||||
wavData := make([]byte, 0, len(wavHeader)+len(audioData))
|
||||
wavData = append(wavData, wavHeader...)
|
||||
wavData = append(wavData, audioData...)
|
||||
audioData = wavData
|
||||
}
|
||||
|
||||
return &dto.SynthesisResult{AudioData: audioData, ReqID: reqID, Format: outputFormat}, nil
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
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
|
||||
MaxModelNameLength = 64
|
||||
MaxRateLimiterEntries = 100000
|
||||
)
|
||||
@@ -0,0 +1,205 @@
|
||||
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()
|
||||
}
|
||||
|
||||
// truncateForLog 用于在日志中安全地展示请求内容(截断避免日志爆炸、控制不可打印字符)
|
||||
func truncateForLog(b []byte, max int) string {
|
||||
if len(b) > max {
|
||||
return string(b[:max]) + fmt.Sprintf("...(truncated, total %d bytes)", len(b))
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
log.Printf("警告: 错误的方法 - 方法=%s 期望=POST 路径=%s 客户端=%s",
|
||||
r.Method, r.URL.Path, middleware.GetClientIP(r))
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
if !middleware.ValidateAPIKey(r) {
|
||||
log.Printf("警告: API Key 鉴权失败 - 路径=%s 客户端=%s 远端=%s",
|
||||
r.URL.Path, middleware.GetClientIP(r), r.RemoteAddr)
|
||||
middleware.SendJSONError(w, http.StatusUnauthorized, "Invalid API key provided.", "invalid_request_error", "invalid_api_key")
|
||||
return
|
||||
}
|
||||
|
||||
if setting.TTSConfigErr != nil {
|
||||
log.Printf("警告: TTS配置未就绪,拒绝请求 - 错误=%v 路径=%s 客户端=%s",
|
||||
setting.TTSConfigErr, r.URL.Path, middleware.GetClientIP(r))
|
||||
middleware.SendJSONError(w, http.StatusServiceUnavailable, "TTS service configuration error. Please check environment variables and restart the service.", "configuration_error", "service_unavailable")
|
||||
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") {
|
||||
log.Printf("警告: 请求体过大 - 路径=%s 客户端=%s 限制=%d字节",
|
||||
r.URL.Path, middleware.GetClientIP(r), common.MaxRequestBodySize)
|
||||
http.Error(w, "Request body too large", http.StatusRequestEntityTooLarge)
|
||||
return
|
||||
}
|
||||
log.Printf("警告: 读取请求体失败 - 路径=%s 客户端=%s 错误=%v",
|
||||
r.URL.Path, middleware.GetClientIP(r), err)
|
||||
http.Error(w, "Failed to read request body", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
var req dto.OpenAITTSRequest
|
||||
if err := json.Unmarshal(body, &req); err != nil {
|
||||
log.Printf("警告: JSON 解析失败 - 路径=%s 客户端=%s 错误=%v body前200字节=%q",
|
||||
r.URL.Path, middleware.GetClientIP(r), err, truncateForLog(body, 200))
|
||||
http.Error(w, "Invalid JSON", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
if req.Model != "" {
|
||||
if len(req.Model) > common.MaxModelNameLength {
|
||||
log.Printf("警告: Model 名过长 - 路径=%s 客户端=%s 长度=%d 限制=%d",
|
||||
r.URL.Path, middleware.GetClientIP(r), len(req.Model), common.MaxModelNameLength)
|
||||
http.Error(w, fmt.Sprintf("Model name too long (max %d characters)", common.MaxModelNameLength), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if strings.ContainsAny(req.Model, "\x00\n\r\t") {
|
||||
log.Printf("警告: Model 名含非法字符 - 路径=%s 客户端=%s model前50字节=%q",
|
||||
r.URL.Path, middleware.GetClientIP(r), truncateForLog([]byte(req.Model), 50))
|
||||
http.Error(w, "Model name contains invalid characters", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if req.Input == "" {
|
||||
log.Printf("警告: input 字段为空 - 路径=%s 客户端=%s", r.URL.Path, middleware.GetClientIP(r))
|
||||
http.Error(w, "Input text is required", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
if len(req.Input) > common.MaxTextLength {
|
||||
log.Printf("警告: input 文本过长 - 路径=%s 客户端=%s 长度=%d 限制=%d",
|
||||
r.URL.Path, middleware.GetClientIP(r), 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
|
||||
}
|
||||
|
||||
// 将 OpenAI response_format 映射为火山 API 支持的格式
|
||||
var requestFormat string
|
||||
if req.ResponseFormat != "" {
|
||||
requestFormat = volcano.MapOpenAIFormat(req.ResponseFormat)
|
||||
}
|
||||
|
||||
ttsStart := time.Now()
|
||||
result, err := volcano.Synthesis(&setting.TTSConfig, volcanoClient, req.Input, speed, req.Voice, requestFormat)
|
||||
duration := time.Since(ttsStart)
|
||||
|
||||
if err != nil {
|
||||
service.GlobalStats.AddRequest(false, duration, err.Error())
|
||||
log.Printf("警告: TTS 合成失败 - 路径=%s 客户端=%s 文本长度=%d 耗时=%v 错误=%v",
|
||||
r.URL.Path, middleware.GetClientIP(r), len(req.Input), duration, err)
|
||||
http.Error(w, "TTS synthesis failed", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
service.GlobalStats.AddRequest(true, duration, "")
|
||||
|
||||
w.Header().Set("Content-Type", volcano.FormatContentType(result.Format))
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
version: '3.8'
|
||||
|
||||
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}
|
||||
- BYTEDANCE_TTS_FORMAT=${BYTEDANCE_TTS_FORMAT:-mp3}
|
||||
- BYTEDANCE_TTS_SAMPLE_RATE=${BYTEDANCE_TTS_SAMPLE_RATE:-24000}
|
||||
- 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
|
||||
@@ -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"`
|
||||
}
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
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
|
||||
Model string // v3 声音复刻/语音大模型 子模型版本,复刻音色必填
|
||||
URL string
|
||||
Timeout time.Duration
|
||||
Format string // 音频编码格式: mp3/ogg_opus/pcm/wav(wav内部用pcm请求再封装header)
|
||||
SampleRate int // 采样率: 8000/16000/22050/24000/32000/44100/48000
|
||||
}
|
||||
|
||||
type SynthesisResult struct {
|
||||
AudioData []byte
|
||||
ReqID string
|
||||
Format string // 实际输出格式,用于设置 Content-Type
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
module bytedance-tts-openai-adapter
|
||||
module github.com/volcano-tts/tts-api
|
||||
|
||||
go 1.19
|
||||
go 1.26
|
||||
|
||||
require (
|
||||
github.com/google/uuid v1.6.0
|
||||
|
||||
+492
@@ -0,0 +1,492 @@
|
||||
<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>TTS 服务监控</title>
|
||||
<script src="https://unpkg.com/vue@3/dist/vue.global.prod.js"></script>
|
||||
<script src="https://unpkg.com/axios/dist/axios.min.js"></script>
|
||||
<style>
|
||||
* {
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
box-sizing: border-box;
|
||||
}
|
||||
body {
|
||||
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, 'Helvetica Neue', Arial, sans-serif;
|
||||
background: linear-gradient(135deg, #1a1a2e 0%, #16213e 100%);
|
||||
min-height: 100vh;
|
||||
color: #e0e0e0;
|
||||
padding: 20px;
|
||||
}
|
||||
#app {
|
||||
max-width: 1200px;
|
||||
margin: 0 auto;
|
||||
}
|
||||
.header {
|
||||
text-align: center;
|
||||
margin-bottom: 30px;
|
||||
}
|
||||
.header h1 {
|
||||
font-size: 2em;
|
||||
background: linear-gradient(90deg, #00d4ff, #7b2ff7);
|
||||
-webkit-background-clip: text;
|
||||
-webkit-text-fill-color: transparent;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.header .version {
|
||||
color: #888;
|
||||
font-size: 0.9em;
|
||||
}
|
||||
.refresh-btn {
|
||||
background: linear-gradient(90deg, #00d4ff, #7b2ff7);
|
||||
border: none;
|
||||
color: white;
|
||||
padding: 10px 24px;
|
||||
border-radius: 8px;
|
||||
cursor: pointer;
|
||||
font-size: 14px;
|
||||
margin-top: 15px;
|
||||
transition: opacity 0.3s;
|
||||
}
|
||||
.refresh-btn:hover {
|
||||
opacity: 0.9;
|
||||
}
|
||||
.refresh-btn:disabled {
|
||||
opacity: 0.5;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
.grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fit, minmax(300px, 1fr));
|
||||
gap: 20px;
|
||||
margin-bottom: 20px;
|
||||
transition: all 0.3s ease;
|
||||
}
|
||||
.card {
|
||||
background: rgba(255, 255, 255, 0.05);
|
||||
border-radius: 16px;
|
||||
padding: 24px;
|
||||
backdrop-filter: blur(10px);
|
||||
border: 1px solid rgba(255, 255, 255, 0.1);
|
||||
transition: all 0.3s ease;
|
||||
}
|
||||
.card-title {
|
||||
font-size: 14px;
|
||||
color: #888;
|
||||
text-transform: uppercase;
|
||||
letter-spacing: 1px;
|
||||
margin-bottom: 16px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
}
|
||||
.card-title .dot {
|
||||
width: 8px;
|
||||
height: 8px;
|
||||
border-radius: 50%;
|
||||
background: #00d4ff;
|
||||
}
|
||||
.card-title .dot.error {
|
||||
background: #ff4757;
|
||||
}
|
||||
.card-title .dot.warning {
|
||||
background: #ffa502;
|
||||
}
|
||||
.stat-value {
|
||||
font-size: 2.5em;
|
||||
font-weight: bold;
|
||||
background: linear-gradient(90deg, #00d4ff, #7b2ff7);
|
||||
-webkit-background-clip: text;
|
||||
-webkit-text-fill-color: transparent;
|
||||
transition: all 0.3s ease;
|
||||
}
|
||||
.stat-label {
|
||||
color: #888;
|
||||
font-size: 14px;
|
||||
margin-top: 5px;
|
||||
}
|
||||
.info-row {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
padding: 12px 0;
|
||||
border-bottom: 1px solid rgba(255, 255, 255, 0.05);
|
||||
}
|
||||
.info-row:last-child {
|
||||
border-bottom: none;
|
||||
}
|
||||
.info-label {
|
||||
color: #888;
|
||||
}
|
||||
.info-value {
|
||||
color: #fff;
|
||||
font-family: 'Monaco', 'Menlo', monospace;
|
||||
}
|
||||
.info-value.success {
|
||||
color: #2ed573;
|
||||
}
|
||||
.info-value.error {
|
||||
color: #ff4757;
|
||||
}
|
||||
.info-value.warning {
|
||||
color: #ffa502;
|
||||
}
|
||||
.chart-container {
|
||||
height: 120px;
|
||||
display: flex;
|
||||
align-items: flex-end;
|
||||
gap: 2px;
|
||||
padding: 10px 0;
|
||||
}
|
||||
.bar {
|
||||
flex: 1;
|
||||
background: linear-gradient(180deg, #7b2ff7, #00d4ff);
|
||||
border-radius: 4px 4px 0 0;
|
||||
min-height: 2px;
|
||||
transition: height 0.3s ease;
|
||||
}
|
||||
.error-list {
|
||||
max-height: 200px;
|
||||
overflow-y: auto;
|
||||
}
|
||||
.error-item {
|
||||
background: rgba(255, 71, 87, 0.1);
|
||||
border-left: 3px solid #ff4757;
|
||||
padding: 10px 12px;
|
||||
margin-bottom: 8px;
|
||||
border-radius: 0 8px 8px 0;
|
||||
font-size: 13px;
|
||||
word-break: break-all;
|
||||
}
|
||||
.error-time {
|
||||
color: #888;
|
||||
font-size: 12px;
|
||||
margin-bottom: 4px;
|
||||
}
|
||||
.loading {
|
||||
text-align: center;
|
||||
padding: 40px;
|
||||
color: #888;
|
||||
}
|
||||
.error-box {
|
||||
background: rgba(255, 71, 87, 0.1);
|
||||
border: 1px solid rgba(255, 71, 87, 0.3);
|
||||
border-radius: 12px;
|
||||
padding: 20px;
|
||||
color: #ff4757;
|
||||
text-align: center;
|
||||
transition: all 0.3s ease;
|
||||
animation: fadeIn 0.3s ease;
|
||||
}
|
||||
@keyframes fadeIn {
|
||||
from { opacity: 0; transform: translateY(-10px); }
|
||||
to { opacity: 1; transform: translateY(0); }
|
||||
}
|
||||
.uptime {
|
||||
font-size: 1.5em;
|
||||
font-weight: bold;
|
||||
color: #2ed573;
|
||||
}
|
||||
.progress-ring {
|
||||
width: 100px;
|
||||
height: 100px;
|
||||
margin: 0 auto;
|
||||
}
|
||||
.progress-ring circle {
|
||||
fill: none;
|
||||
stroke-width: 8;
|
||||
}
|
||||
.progress-ring .bg {
|
||||
stroke: rgba(255, 255, 255, 0.1);
|
||||
}
|
||||
.progress-ring .progress {
|
||||
stroke: url(#gradient);
|
||||
stroke-linecap: round;
|
||||
transform: rotate(-90deg);
|
||||
transform-origin: center;
|
||||
transition: stroke-dashoffset 0.5s ease;
|
||||
}
|
||||
.progress-text {
|
||||
position: absolute;
|
||||
top: 50%;
|
||||
left: 50%;
|
||||
transform: translate(-50%, -50%);
|
||||
text-align: center;
|
||||
}
|
||||
.memory-stat {
|
||||
display: flex;
|
||||
justify-content: space-around;
|
||||
text-align: center;
|
||||
}
|
||||
.memory-stat .value {
|
||||
font-size: 1.2em;
|
||||
font-weight: bold;
|
||||
color: #00d4ff;
|
||||
}
|
||||
.memory-stat .label {
|
||||
font-size: 12px;
|
||||
color: #888;
|
||||
margin-top: 4px;
|
||||
}
|
||||
.no-errors {
|
||||
text-align: center;
|
||||
color: #2ed573;
|
||||
padding: 20px;
|
||||
}
|
||||
::-webkit-scrollbar {
|
||||
width: 6px;
|
||||
}
|
||||
::-webkit-scrollbar-track {
|
||||
background: rgba(255, 255, 255, 0.05);
|
||||
}
|
||||
::-webkit-scrollbar-thumb {
|
||||
background: rgba(255, 255, 255, 0.2);
|
||||
border-radius: 3px;
|
||||
}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div id="app">
|
||||
<div class="header">
|
||||
<h1>TTS 服务监控</h1>
|
||||
<div class="version">{{ healthData.service || 'Loading...' }} - {{ healthData.version || '' }}</div>
|
||||
<button class="refresh-btn" @click="fetchHealth(true)" :disabled="loading">
|
||||
{{ loading ? '刷新中...' : '刷新数据' }}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div v-if="error" class="error-box">
|
||||
{{ error }}
|
||||
</div>
|
||||
|
||||
<div v-if="healthData.status">
|
||||
<div class="grid">
|
||||
<div class="card">
|
||||
<div class="card-title">
|
||||
<span class="dot" :class="{ error: healthData.status !== 'ok' }"></span>
|
||||
服务状态
|
||||
</div>
|
||||
<div class="stat-value">{{ healthData.status === 'ok' ? '正常运行' : '配置错误' }}</div>
|
||||
<div class="stat-label">运行时长: {{ formatUptime(healthData.uptime) }}</div>
|
||||
</div>
|
||||
|
||||
<div class="card">
|
||||
<div class="card-title">
|
||||
<span class="dot"></span>
|
||||
请求统计
|
||||
</div>
|
||||
<div style="display: flex; gap: 30px;">
|
||||
<div>
|
||||
<div class="stat-value">{{ healthData.api_stats?.total_requests || 0 }}</div>
|
||||
<div class="stat-label">总请求数</div>
|
||||
</div>
|
||||
<div>
|
||||
<div class="stat-value" style="color: #2ed573">{{ healthData.api_stats?.successful_requests || 0 }}</div>
|
||||
<div class="stat-label">成功</div>
|
||||
</div>
|
||||
<div>
|
||||
<div class="stat-value" style="color: #ff4757">{{ healthData.api_stats?.failed_requests || 0 }}</div>
|
||||
<div class="stat-label">失败</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="card">
|
||||
<div class="card-title">
|
||||
<span class="dot" :class="{ warning: errorRate > 10 }"></span>
|
||||
错误率
|
||||
</div>
|
||||
<div class="stat-value" :style="{ color: errorRate > 10 ? '#ff4757' : '#2ed573' }">
|
||||
{{ healthData.api_stats?.error_rate_percent || '0' }}%
|
||||
</div>
|
||||
<div class="stat-label">平均响应: {{ healthData.api_stats?.avg_response_time_ms || '0' }} ms</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="grid">
|
||||
<div class="card">
|
||||
<div class="card-title">
|
||||
<span class="dot"></span>
|
||||
配置状态
|
||||
</div>
|
||||
<div class="info-row">
|
||||
<span class="info-label">环境变量</span>
|
||||
<span class="info-value" :class="healthData.config_status?.all_required_vars_set ? 'success' : 'error'">
|
||||
{{ healthData.config_status?.all_required_vars_set ? '已配置' : '未配置' }}
|
||||
</span>
|
||||
</div>
|
||||
<div class="info-row">
|
||||
<span class="info-label">配置状态</span>
|
||||
<span class="info-value" :class="healthData.config_status?.config_error ? 'error' : 'success'">
|
||||
{{ healthData.config_status?.config_error ? '异常' : '正常' }}
|
||||
</span>
|
||||
</div>
|
||||
<div class="info-row">
|
||||
<span class="info-label">启动时间</span>
|
||||
<span class="info-value">{{ healthData.start_time || '-' }}</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="card">
|
||||
<div class="card-title">
|
||||
<span class="dot"></span>
|
||||
内存使用
|
||||
</div>
|
||||
<div class="memory-stat">
|
||||
<div>
|
||||
<div class="value">{{ formatBytes(healthData.memory?.heap_alloc) }}</div>
|
||||
<div class="label">Heap Alloc</div>
|
||||
</div>
|
||||
<div>
|
||||
<div class="value">{{ formatBytes(healthData.memory?.heap_inuse) }}</div>
|
||||
<div class="label">Heap Inuse</div>
|
||||
</div>
|
||||
<div>
|
||||
<div class="value">{{ healthData.memory?.goroutines || 0 }}</div>
|
||||
<div class="label">Goroutines</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="card">
|
||||
<div class="card-title">
|
||||
<span class="dot"></span>
|
||||
响应时间趋势
|
||||
</div>
|
||||
<div class="chart-container">
|
||||
<div v-for="(time, index) in chartData" :key="index" class="bar"
|
||||
:style="{ height: Math.max(2, (time / maxResponseTime) * 100) + '%' }"
|
||||
:title="time.toFixed(1) + 'ms'">
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="card">
|
||||
<div class="card-title">
|
||||
<span class="dot" :class="{ error: recentErrorsCount > 0 }"></span>
|
||||
错误记录 ({{ recentErrorsCount }})
|
||||
</div>
|
||||
<div v-if="recentErrorsCount === 0" class="no-errors">
|
||||
暂无错误记录
|
||||
</div>
|
||||
<div v-else class="error-list">
|
||||
<div class="error-item">
|
||||
<div class="error-time">检测到 {{ recentErrorsCount }} 条最近错误,详细信息请查看服务器日志</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div v-if="loading && !healthData.status" class="loading">
|
||||
加载中...
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<script>
|
||||
const { createApp, ref, computed, onMounted } = Vue;
|
||||
|
||||
createApp({
|
||||
setup() {
|
||||
const healthData = ref({});
|
||||
const loading = ref(false);
|
||||
const error = ref(null);
|
||||
const isAutoRefresh = ref(false);
|
||||
|
||||
const deepUpdate = (target, source) => {
|
||||
for (const key of Object.keys(source)) {
|
||||
if (source[key] && typeof source[key] === 'object' && !Array.isArray(source[key])) {
|
||||
if (!target[key]) target[key] = {};
|
||||
deepUpdate(target[key], source[key]);
|
||||
} else {
|
||||
target[key] = source[key];
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
const fetchHealth = async (showLoading = false) => {
|
||||
if (showLoading) {
|
||||
loading.value = true;
|
||||
}
|
||||
const oldError = error.value;
|
||||
try {
|
||||
const response = await axios.get('/health');
|
||||
deepUpdate(healthData.value, response.data);
|
||||
if (response.data.config_status?.all_required_vars_set !== false) {
|
||||
error.value = null;
|
||||
}
|
||||
} catch (e) {
|
||||
if (e.response && e.response.data) {
|
||||
deepUpdate(healthData.value, e.response.data);
|
||||
if (e.response.status === 503) {
|
||||
error.value = '服务配置异常,请检查环境变量配置';
|
||||
} else {
|
||||
error.value = '服务异常: ' + (e.message || '未知错误');
|
||||
}
|
||||
} else {
|
||||
error.value = '无法获取服务状态: ' + (e.message || '未知错误');
|
||||
}
|
||||
} finally {
|
||||
loading.value = false;
|
||||
}
|
||||
};
|
||||
|
||||
const errorRate = computed(() => {
|
||||
return parseFloat(healthData.value.api_stats?.error_rate_percent || 0);
|
||||
});
|
||||
|
||||
const recentErrorsCount = computed(() => {
|
||||
return healthData.value.errors?.recent_errors_count || 0;
|
||||
});
|
||||
|
||||
const chartData = computed(() => {
|
||||
return healthData.value.api_stats?.recent_response_times_ms || [];
|
||||
});
|
||||
|
||||
const maxResponseTime = computed(() => {
|
||||
const times = chartData.value;
|
||||
if (times.length === 0) return 100;
|
||||
return Math.max(...times, 100);
|
||||
});
|
||||
|
||||
const formatBytes = (bytes) => {
|
||||
if (!bytes) return '0 B';
|
||||
const k = 1024;
|
||||
const sizes = ['B', 'KB', 'MB', 'GB'];
|
||||
const i = Math.floor(Math.log(bytes) / Math.log(k));
|
||||
return (bytes / Math.pow(k, i)).toFixed(1) + ' ' + sizes[i];
|
||||
};
|
||||
|
||||
const formatUptime = (seconds) => {
|
||||
if (!seconds) return '-';
|
||||
const s = parseInt(seconds);
|
||||
const d = Math.floor(s / 86400);
|
||||
const h = Math.floor((s % 86400) / 3600);
|
||||
const m = Math.floor((s % 3600) / 60);
|
||||
if (d > 0) return `${d}天 ${h}小时`;
|
||||
if (h > 0) return `${h}小时 ${m}分钟`;
|
||||
return `${m}分钟`;
|
||||
};
|
||||
|
||||
onMounted(() => {
|
||||
fetchHealth(true);
|
||||
setInterval(() => fetchHealth(false), 10000);
|
||||
});
|
||||
|
||||
return {
|
||||
healthData,
|
||||
loading,
|
||||
error,
|
||||
fetchHealth,
|
||||
errorRate,
|
||||
recentErrorsCount,
|
||||
chartData,
|
||||
maxResponseTime,
|
||||
formatBytes,
|
||||
formatUptime
|
||||
};
|
||||
}
|
||||
}).mount('#app');
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -0,0 +1,74 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"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] ")
|
||||
|
||||
// 所有环境变量读取在 setting 包内集中完成,业务模块只读全局 Config。
|
||||
setting.InitAllConfigs()
|
||||
|
||||
// 兼容旧调用顺序:rate limiter / 静态文件 / stats / controller 的初始化保持独立。
|
||||
middleware.InitRateLimiter()
|
||||
setting.CheckStaticFiles()
|
||||
service.InitStats()
|
||||
controller.InitController()
|
||||
|
||||
// 启动期一次性打印所有 Config 状态,便于运维核对。
|
||||
// (必填项缺失的明确警告由 LogStartupSummary 自身负责,避免重复打印。)
|
||||
setting.LogStartupSummary()
|
||||
|
||||
controller.SetStartTime(time.Now())
|
||||
|
||||
r := router.Setup()
|
||||
|
||||
server := &http.Server{
|
||||
Addr: ":" + setting.Server.Port,
|
||||
Handler: middleware.CORS(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", setting.Server.Port)
|
||||
log.Printf("OpenAI TTS endpoint: http://localhost:%s/v1/audio/speech", setting.Server.Port)
|
||||
log.Printf("Health check: http://localhost:%s/health", setting.Server.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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"crypto/subtle"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/volcano-tts/tts-api/setting"
|
||||
)
|
||||
|
||||
// InitAPIKeys 已在 setting.InitAuthConfig 中完成,这里保留为 no-op 以维持现有调用顺序。
|
||||
// 实际鉴权逻辑直接读 setting.Auth.APIKeys。
|
||||
func InitAPIKeys() {
|
||||
// 配置由 setting 包统一加载,日志也由 setting.LogStartupSummary 输出。
|
||||
_ = setting.Auth
|
||||
}
|
||||
|
||||
func ValidateAPIKey(r *http.Request) bool {
|
||||
if len(setting.Auth.APIKeys) == 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 setting.Auth.APIKeys {
|
||||
if subtle.ConstantTimeCompare([]byte(token), []byte(validKey)) == 1 {
|
||||
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,
|
||||
},
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"log"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/volcano-tts/tts-api/setting"
|
||||
)
|
||||
|
||||
var (
|
||||
corsMaxAgeHeader = "86400"
|
||||
)
|
||||
|
||||
// InitCORSConfig 已在 setting.InitCORSConfig 中完成,这里保留为 no-op 以维持现有调用顺序。
|
||||
// 实际 CORS 匹配逻辑直接读 setting.CORS.Origins / setting.CORS.AllowAll。
|
||||
func InitCORSConfig() {
|
||||
// 配置由 setting 包统一加载,日志也由 setting.LogStartupSummary 输出。
|
||||
_ = setting.CORS
|
||||
}
|
||||
|
||||
func isValidOrigin(origin string) bool {
|
||||
if origin == "" || origin == "null" || origin == "nil" {
|
||||
return false
|
||||
}
|
||||
lowerOrigin := strings.ToLower(origin)
|
||||
if !strings.HasPrefix(lowerOrigin, "http://") && !strings.HasPrefix(lowerOrigin, "https://") {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func matchOrigin(origin string) (string, bool) {
|
||||
if !isValidOrigin(origin) {
|
||||
return "", false
|
||||
}
|
||||
if setting.CORS.AllowAll {
|
||||
return "*", true
|
||||
}
|
||||
normalized := strings.ToLower(strings.TrimRight(strings.TrimSpace(origin), "/"))
|
||||
for _, allowed := range setting.CORS.Origins {
|
||||
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")
|
||||
|
||||
// 无 Origin 头:非跨域请求,跳过 CORS 处理
|
||||
if origin == "" {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
// 有 Origin 头时,响应必须携带 Vary: Origin 防止 CDN 缓存污染
|
||||
vary := w.Header().Get("Vary")
|
||||
if vary == "" {
|
||||
w.Header().Set("Vary", "Origin")
|
||||
} else if !strings.Contains(vary, "Origin") {
|
||||
w.Header().Set("Vary", vary+", Origin")
|
||||
}
|
||||
|
||||
isPreflight := r.Method == http.MethodOptions
|
||||
|
||||
allowOrigin, matched := matchOrigin(origin)
|
||||
if !matched {
|
||||
// Origin 不在白名单:拒绝请求(预检和非预检均拒绝),
|
||||
// 防止不匹配的请求穿透到后端浪费 TTS 资源
|
||||
log.Printf("CORS拦截: 来源=%q 路径=%s 方法=%s 客户端=%s",
|
||||
origin, r.URL.Path, r.Method, GetClientIP(r))
|
||||
w.WriteHeader(http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
// Origin 匹配:设置 CORS 响应头
|
||||
w.Header().Set("Access-Control-Allow-Origin", allowOrigin)
|
||||
w.Header().Set("Access-Control-Allow-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")
|
||||
}
|
||||
|
||||
// 预检请求:直接返回 204,不进入内层中间件链,
|
||||
// 避免消耗速率限制配额和并发槽位
|
||||
if isPreflight {
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
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) {
|
||||
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)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"log"
|
||||
"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
|
||||
}
|
||||
}
|
||||
|
||||
if len(rl.requests) > common.MaxRateLimiterEntries {
|
||||
log.Printf("警告: 限流器条目数 %d 超过上限 %d,触发强制清理", len(rl.requests), common.MaxRateLimiterEntries)
|
||||
for k := range rl.requests {
|
||||
if len(rl.requests) <= common.MaxRateLimiterEntries/2 {
|
||||
break
|
||||
}
|
||||
delete(rl.requests, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 私有网络 CIDR 范围:仅在直连来源属于这些范围时才信任代理头
|
||||
var privateCIDRs []*net.IPNet
|
||||
|
||||
func init() {
|
||||
for _, cidr := range []string{
|
||||
"10.0.0.0/8",
|
||||
"172.16.0.0/12",
|
||||
"192.168.0.0/16",
|
||||
"127.0.0.0/8",
|
||||
"169.254.0.0/16",
|
||||
"::1/128",
|
||||
"fc00::/7",
|
||||
"fe80::/10",
|
||||
} {
|
||||
_, ipNet, _ := net.ParseCIDR(cidr)
|
||||
privateCIDRs = append(privateCIDRs, ipNet)
|
||||
}
|
||||
}
|
||||
|
||||
func isPrivateIP(ipStr string) bool {
|
||||
ip := net.ParseIP(ipStr)
|
||||
if ip == nil {
|
||||
return false
|
||||
}
|
||||
if ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() {
|
||||
return true
|
||||
}
|
||||
for _, cidr := range privateCIDRs {
|
||||
if cidr.Contains(ip) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// GetClientIP 提取客户端真实 IP。
|
||||
// 仅当直连来源为私有网络(本地代理、Docker 网桥等)时才信任 X-Forwarded-For / X-Real-IP,
|
||||
// 防止公网直连场景下攻击者伪造代理头绕过速率限制。
|
||||
func GetClientIP(r *http.Request) string {
|
||||
directIP, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
directIP = r.RemoteAddr
|
||||
}
|
||||
|
||||
if isPrivateIP(directIP) {
|
||||
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
|
||||
ip := strings.TrimSpace(strings.Split(xff, ",")[0])
|
||||
if net.ParseIP(ip) != nil {
|
||||
return ip
|
||||
}
|
||||
}
|
||||
if xri := strings.TrimSpace(r.Header.Get("X-Real-IP")); xri != "" {
|
||||
if net.ParseIP(xri) != nil {
|
||||
return xri
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return directIP
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"log"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
func RateLimit(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
clientIP := GetClientIP(r)
|
||||
if !GlobalRateLimiter.Allow(clientIP) {
|
||||
log.Printf("警告: 已超过IP速率限制,拒绝请求 - 客户端IP: %s", clientIP)
|
||||
SendJSONError(w, http.StatusTooManyRequests, "Rate limit exceeded. Please try again later.", "rate_limit_error", "rate_limit_exceeded")
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func ConcurrencyLimit(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
select {
|
||||
case ConcurrencySem <- struct{}{}:
|
||||
defer func() { <-ConcurrencySem }()
|
||||
next.ServeHTTP(w, r)
|
||||
default:
|
||||
log.Printf("警告: 已达到最大并发请求数限制,拒绝请求 - 客户端IP: %s", GetClientIP(r))
|
||||
SendJSONError(w, http.StatusServiceUnavailable, "Server is busy, maximum concurrent requests reached. Please try again later.", "concurrency_limit_error", "max_concurrent_requests")
|
||||
return
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func SecurityHeaders(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
w.Header().Set("X-Frame-Options", "DENY")
|
||||
w.Header().Set("X-XSS-Protection", "1; mode=block")
|
||||
w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin")
|
||||
|
||||
if strings.HasPrefix(r.URL.Path, "/v1/") || r.URL.Path == "/health" {
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
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.SecurityHeaders)
|
||||
r.Use(middleware.RateLimit)
|
||||
r.Use(middleware.ConcurrencyLimit)
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"runtime"
|
||||
"strings"
|
||||
"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
|
||||
responseTimesCount int
|
||||
lastErrors []string
|
||||
errorsIndex int
|
||||
errorsCount 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 s.responseTimesCount < common.MaxResponseTimes {
|
||||
s.responseTimesCount++
|
||||
}
|
||||
|
||||
if success {
|
||||
s.successfulRequests++
|
||||
} else {
|
||||
s.failedRequests++
|
||||
if errMsg != "" {
|
||||
now := time.Now().Format(time.RFC3339)
|
||||
|
||||
// 去重:如果最近一条错误的消息内容相同,仅更新时间戳
|
||||
if s.errorsCount > 0 {
|
||||
lastIdx := (s.errorsIndex - 1 + common.MaxErrors) % common.MaxErrors
|
||||
lastEntry := s.lastErrors[lastIdx]
|
||||
if sepIdx := strings.Index(lastEntry, ": "); sepIdx != -1 {
|
||||
if lastEntry[sepIdx+2:] == errMsg {
|
||||
s.lastErrors[lastIdx] = now + ": " + errMsg
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
s.lastErrors[s.errorsIndex] = now + ": " + errMsg
|
||||
s.errorsIndex = (s.errorsIndex + 1) % common.MaxErrors
|
||||
if s.errorsCount < common.MaxErrors {
|
||||
s.errorsCount++
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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, s.responseTimesCount)
|
||||
if s.responseTimesCount > 0 {
|
||||
start := 0
|
||||
if s.responseTimesCount == common.MaxResponseTimes {
|
||||
start = s.responseTimesIndex
|
||||
}
|
||||
for i := 0; i < s.responseTimesCount; i++ {
|
||||
idx := (start + i) % common.MaxResponseTimes
|
||||
recentResponseTimes = append(recentResponseTimes, s.recentResponseTimes[idx])
|
||||
}
|
||||
}
|
||||
|
||||
// 按时间顺序(从旧到新)遍历错误环形缓冲区
|
||||
lastErrors = make([]string, 0, s.errorsCount)
|
||||
if s.errorsCount > 0 {
|
||||
start := 0
|
||||
if s.errorsCount == common.MaxErrors {
|
||||
start = s.errorsIndex
|
||||
}
|
||||
for i := 0; i < s.errorsCount; i++ {
|
||||
idx := (start + i) % common.MaxErrors
|
||||
lastErrors = append(lastErrors, s.lastErrors[idx])
|
||||
}
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
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(),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,287 @@
|
||||
package setting
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/volcano-tts/tts-api/common"
|
||||
"github.com/volcano-tts/tts-api/dto"
|
||||
)
|
||||
|
||||
// 全部环境变量读取的单一入口:其它包不允许直接 os.Getenv,只读这里的全局 Config。
|
||||
|
||||
// TTSConfig 上游火山 TTS 配置(由 InitTTSConfig 填充)。
|
||||
var (
|
||||
TTSConfig dto.ByteDanceTTSConfig
|
||||
TTSConfigErr error
|
||||
)
|
||||
|
||||
// AuthConfig OpenAI 兼容接口的客户端 API Key 鉴权配置。
|
||||
type AuthConfig struct {
|
||||
APIKeys []string
|
||||
}
|
||||
|
||||
var Auth AuthConfig
|
||||
|
||||
// CORSConfig 跨域白名单配置。
|
||||
type CORSConfig struct {
|
||||
Origins []string
|
||||
AllowAll bool
|
||||
}
|
||||
|
||||
var CORS CORSConfig
|
||||
|
||||
// ServerConfig HTTP 服务监听配置。
|
||||
type ServerConfig struct {
|
||||
Port string
|
||||
}
|
||||
|
||||
var Server ServerConfig
|
||||
|
||||
// InitAllConfigs 集中初始化所有配置,启动期调用一次。
|
||||
// 返回 TTSConfigErr(火山 TTS 必填项缺失时为非 nil);其它 Config 缺失时不返回 error,
|
||||
// 各自有合理兜底(Auth 放行 / CORS 拒绝跨域 / Server 默认 8080)。
|
||||
func InitAllConfigs() {
|
||||
InitServerConfig()
|
||||
InitAuthConfig()
|
||||
InitCORSConfig()
|
||||
TTSConfigErr = InitTTSConfig()
|
||||
}
|
||||
|
||||
// InitServerConfig 读取 PORT,缺省 common.DefaultPort。
|
||||
func InitServerConfig() {
|
||||
Server.Port = os.Getenv("PORT")
|
||||
if Server.Port == "" {
|
||||
Server.Port = common.DefaultPort
|
||||
}
|
||||
}
|
||||
|
||||
// InitAuthConfig 读取 OPENAI_TTS_API_KEY,支持逗号分隔多个 key。
|
||||
// 留空时 Auth.APIKeys 为空,ValidateAPIKey 会放行所有请求。
|
||||
func InitAuthConfig() {
|
||||
raw := os.Getenv("OPENAI_TTS_API_KEY")
|
||||
if raw == "" {
|
||||
Auth.APIKeys = nil
|
||||
return
|
||||
}
|
||||
parts := strings.Split(raw, ",")
|
||||
keys := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
k := strings.TrimSpace(p)
|
||||
if k != "" {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
}
|
||||
Auth.APIKeys = keys
|
||||
}
|
||||
|
||||
// InitCORSConfig 读取 ALLOWED_ORIGINS,按逗号分隔;支持 * 通配(AllowAll=true)。
|
||||
// 留空时 CORS.Origins 为空,跨域请求会被拒绝。
|
||||
func InitCORSConfig() {
|
||||
raw := os.Getenv("ALLOWED_ORIGINS")
|
||||
CORS.Origins = nil
|
||||
CORS.AllowAll = false
|
||||
if raw == "" {
|
||||
return
|
||||
}
|
||||
for _, p := range strings.Split(raw, ",") {
|
||||
o := strings.TrimSpace(p)
|
||||
if o == "" {
|
||||
continue
|
||||
}
|
||||
if o == "*" {
|
||||
CORS.AllowAll = true
|
||||
continue
|
||||
}
|
||||
CORS.Origins = append(CORS.Origins, normalizeOrigin(o))
|
||||
}
|
||||
}
|
||||
|
||||
// normalizeOrigin 复制自原 middleware/cors.go:小写 + 去尾斜杠。
|
||||
func normalizeOrigin(origin string) string {
|
||||
origin = strings.TrimSpace(origin)
|
||||
origin = strings.TrimRight(origin, "/")
|
||||
return strings.ToLower(origin)
|
||||
}
|
||||
|
||||
// InitTTSConfig 读取火山 TTS 必填和可选配置,填充 TTSConfig。
|
||||
// 必填项缺失时返回 error,服务可继续运行但 TTS 功能不可用。
|
||||
func InitTTSConfig() error {
|
||||
apiKey := os.Getenv("BYTEDANCE_TTS_API_KEY")
|
||||
resourceId := os.Getenv("BYTEDANCE_TTS_RESOURCE_ID")
|
||||
speaker := os.Getenv("BYTEDANCE_TTS_SPEAKER")
|
||||
model := os.Getenv("BYTEDANCE_TTS_MODEL")
|
||||
if model == "" {
|
||||
model = "seed-tts-2.0-standard" // 文档默认值 复刻音色可设为 seed-tts-2.0-expressive
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
// 音频格式,默认 mp3(文档默认值,流式场景中 wav 会多次返回 header,不推荐)
|
||||
format := os.Getenv("BYTEDANCE_TTS_FORMAT")
|
||||
if format == "" {
|
||||
format = "mp3"
|
||||
}
|
||||
|
||||
// 采样率,默认 24000
|
||||
sampleRate := 24000
|
||||
if srStr := os.Getenv("BYTEDANCE_TTS_SAMPLE_RATE"); srStr != "" {
|
||||
if sr, err := fmt.Sscanf(srStr, "%d", &sampleRate); err != nil || sr != 1 {
|
||||
log.Printf("无效的采样率设置 '%s',使用默认值: 24000", srStr)
|
||||
sampleRate = 24000
|
||||
}
|
||||
validRates := map[int]bool{8000: true, 16000: true, 22050: true, 24000: true, 32000: true, 44100: true, 48000: true}
|
||||
if !validRates[sampleRate] {
|
||||
log.Printf("不支持的采样率 %d,使用默认值: 24000", sampleRate)
|
||||
sampleRate = 24000
|
||||
}
|
||||
}
|
||||
|
||||
TTSConfig = dto.ByteDanceTTSConfig{
|
||||
ApiKey: apiKey,
|
||||
ResourceId: resourceId,
|
||||
Speaker: speaker,
|
||||
Model: model,
|
||||
URL: url,
|
||||
Timeout: timeout,
|
||||
Format: format,
|
||||
SampleRate: sampleRate,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// LogStartupSummary 在启动期打印所有 Config 的最终状态。
|
||||
// 调用时机:InitAllConfigs 之后,ListenAndServe 之前。
|
||||
// 必填项逐项输出,失败分支明确告知"v1/audio/speech 路由将 500"。
|
||||
func LogStartupSummary() {
|
||||
log.Printf("=== 环境配置汇总 ===")
|
||||
log.Printf("服务端口: %s", Server.Port)
|
||||
|
||||
if len(Auth.APIKeys) == 0 {
|
||||
log.Printf("OPENAI_TTS_API_KEY: 未设置(所有请求无需鉴权)")
|
||||
} else {
|
||||
log.Printf("OPENAI_TTS_API_KEY: 已设置 %d 个有效密钥", len(Auth.APIKeys))
|
||||
}
|
||||
|
||||
if CORS.AllowAll {
|
||||
log.Printf("ALLOWED_ORIGINS: *(允许所有跨域,不可与凭据共用)")
|
||||
} else if len(CORS.Origins) == 0 {
|
||||
log.Printf("ALLOWED_ORIGINS: 未设置(跨域请求将被拒绝)")
|
||||
} else {
|
||||
log.Printf("ALLOWED_ORIGINS: 已配置 %d 个允许的跨域来源白名单", len(CORS.Origins))
|
||||
}
|
||||
|
||||
// 火山 TTS 必填项逐项状态:缺则 ❌,有则 ✓(API Key 脱敏,仅显示头尾各 4 字符)
|
||||
log.Printf("火山 TTS 必填项状态:")
|
||||
type ttsCheck struct {
|
||||
name string
|
||||
value string
|
||||
ok bool
|
||||
}
|
||||
checks := []ttsCheck{
|
||||
{"BYTEDANCE_TTS_API_KEY", maskAPIKey(TTSConfig.ApiKey), TTSConfig.ApiKey != ""},
|
||||
{"BYTEDANCE_TTS_RESOURCE_ID", TTSConfig.ResourceId, TTSConfig.ResourceId != ""},
|
||||
{"BYTEDANCE_TTS_SPEAKER", TTSConfig.Speaker, TTSConfig.Speaker != ""},
|
||||
}
|
||||
missingCount := 0
|
||||
for _, c := range checks {
|
||||
mark := "✓"
|
||||
if !c.ok {
|
||||
mark = "❌"
|
||||
missingCount++
|
||||
}
|
||||
val := c.value
|
||||
if val == "" {
|
||||
val = "(未设置)"
|
||||
}
|
||||
log.Printf(" %s %s: %s", mark, c.name, val)
|
||||
}
|
||||
|
||||
if TTSConfigErr != nil {
|
||||
log.Printf("火山 TTS 整体: 初始化失败,%d 个必填项缺失,/v1/audio/speech 路由将全部返回 500", missingCount)
|
||||
} else {
|
||||
log.Printf("火山 TTS 可选项: model=%s, format=%s, sample_rate=%d, timeout=%v",
|
||||
TTSConfig.Model, TTSConfig.Format, TTSConfig.SampleRate, TTSConfig.Timeout)
|
||||
log.Printf("火山 TTS 整体: 初始化成功")
|
||||
}
|
||||
}
|
||||
|
||||
// maskAPIKey 对 API Key 脱敏,显示头 4 / 尾 4 字符,中间 * 号代替。
|
||||
// 短于等于 8 字符整体掩为 ****,空串原样返回。
|
||||
func maskAPIKey(key string) string {
|
||||
if key == "" {
|
||||
return ""
|
||||
}
|
||||
if len(key) <= 8 {
|
||||
return "****"
|
||||
}
|
||||
return key[:4] + "****" + key[len(key)-4:]
|
||||
}
|
||||
|
||||
// CheckEnvironmentVariables 返回环境变量状态,供 /health 端点使用。
|
||||
// 不再直接 os.Getenv,改为读已初始化的全局 Config(单一数据源)。
|
||||
func CheckEnvironmentVariables() map[string]interface{} {
|
||||
requiredVars := map[string]bool{
|
||||
"BYTEDANCE_TTS_API_KEY": TTSConfig.ApiKey != "",
|
||||
"BYTEDANCE_TTS_RESOURCE_ID": TTSConfig.ResourceId != "",
|
||||
"BYTEDANCE_TTS_SPEAKER": TTSConfig.Speaker != "",
|
||||
}
|
||||
|
||||
missingVars := []string{}
|
||||
for varName, isSet := range requiredVars {
|
||||
if !isSet {
|
||||
missingVars = append(missingVars, varName)
|
||||
}
|
||||
}
|
||||
|
||||
optionalVars := map[string]bool{
|
||||
|
||||
"BYTEDANCE_TTS_MODEL": TTSConfig.Model != "" && TTSConfig.Model != "seed-tts-2.0-standard",
|
||||
"BYTEDANCE_TTS_FORMAT": TTSConfig.Format != "" && TTSConfig.Format != "mp3",
|
||||
"BYTEDANCE_TTS_SAMPLE_RATE": TTSConfig.SampleRate != 24000,
|
||||
"OPENAI_TTS_API_KEY": len(Auth.APIKeys) > 0,
|
||||
"ALLOWED_ORIGINS": CORS.AllowAll || len(CORS.Origins) > 0,
|
||||
"PORT": Server.Port != common.DefaultPort,
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"all_required_vars_set": len(missingVars) == 0,
|
||||
"missing_required_vars": missingVars,
|
||||
"required_vars_set": requiredVars,
|
||||
"optional_vars_set": optionalVars,
|
||||
}
|
||||
}
|
||||
|
||||
// CheckStaticFiles 静态文件存在性检查,/dashboard 路由需要 health.html。
|
||||
func CheckStaticFiles() {
|
||||
if _, err := os.Stat("health.html"); os.IsNotExist(err) {
|
||||
log.Println("警告: health.html 不存在,/dashboard 路由将返回 404")
|
||||
}
|
||||
}
|
||||
-778
@@ -1,778 +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: 10,
|
||||
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, _ := io.ReadAll(resp.Body)
|
||||
log.Printf("TTS service error: status=%d, body=%s", resp.StatusCode, string(body))
|
||||
return nil, fmt.Errorf("TTS service error: status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
var audioData []byte
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
scanner.Buffer(make([]byte, 1024*1024), 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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(http.MaxBytesReader(w, r.Body, MAX_REQUEST_BODY_SIZE))
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "request body too large") {
|
||||
http.Error(w, "Request body too large", http.StatusRequestEntityTooLarge)
|
||||
} else {
|
||||
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,
|
||||
"config_error_message": fmt.Sprintf("%v", ttsConfigErr),
|
||||
},
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
func corsMiddleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
|
||||
w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization")
|
||||
|
||||
if r.Method == http.MethodOptions {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func main() {
|
||||
startTime = time.Now()
|
||||
|
||||
log.SetFlags(log.LstdFlags | log.Lshortfile)
|
||||
log.SetPrefix("[TTS-Server] ")
|
||||
|
||||
initAPIKeys()
|
||||
|
||||
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("/", func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, "/health", http.StatusFound)
|
||||
}).Methods("GET")
|
||||
|
||||
port := os.Getenv("PORT")
|
||||
if port == "" {
|
||||
port = DEFAULT_PORT
|
||||
}
|
||||
|
||||
server := &http.Server{
|
||||
Addr: ":" + port,
|
||||
Handler: router,
|
||||
ReadTimeout: 15 * time.Second,
|
||||
WriteTimeout: 15 * 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: %s", ttsConfig.URL)
|
||||
log.Printf("Resource ID: %s", ttsConfig.ResourceId)
|
||||
log.Printf("Speaker: %s", ttsConfig.Speaker)
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user