diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..45b0f8d --- /dev/null +++ b/.dockerignore @@ -0,0 +1,9 @@ +*.exe +*.md +.env +.env.example +.git +.gitignore +tts_api_architecture.html +代码审查报告.md +fix_list.md diff --git a/.env.example b/.env.example index a1e9875..a65dfb4 100644 --- a/.env.example +++ b/.env.example @@ -1,4 +1,4 @@ -# ByteDance TTS v3 API 配置示例 +# 字节火山引擎 TTS v3 API 配置示例 # 将此文件复制为 .env 并填入实际配置 # ========================================== @@ -8,32 +8,51 @@ # 火山引擎新版控制台获取的 API Key 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还是2.0模型) +# 复刻 2.0 音色(seed-icl-2.0) +BYTEDANCE_TTS_RESOURCE_ID=seed-icl-2.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 # ========================================== # 可选的环境变量 # ========================================== -# 请求超时时间,默认30秒 +# 单次合成超时,默认30s BYTEDANCE_TTS_TIMEOUT=30s -# OpenAI兼容接口的API密钥(可选) -# 配置后,客户端请求需要携带 Authorization: Bearer +# 上游实际请求的音频格式:mp3 / pcm / ogg_opus +# 客户端要求 wav 时,内部自动转 pcm 上游 + 本地拼 WAV 头 +BYTEDANCE_TTS_FORMAT=mp3 + +# 上游采样率:8000/16000/22050/24000/32000/44100/48000 +BYTEDANCE_TTS_SAMPLE_RATE=24000 + +# MP3 比特率(可选),仅 MP3 生效 +# BYTEDANCE_TTS_BIT_RATE=128000 + +# 复刻 2.0 子模型(可选),留空则使用控制台默认值 +# seed-tts-2.0-standard:标准版,延时更优 +# seed-tts-2.0-expressive:表现力增强版,支持 QA / Cot +# BYTEDANCE_TTS_MODEL=seed-tts-2.0-standard + +# 复刻 2.0 模型类型(可选,推荐显式指定) +# 4 = ICL V2,5 = ICL V3 +# BYTEDANCE_TTS_MODEL_TYPE=4 + +# 非中文/英文合成时指定语种(可选) +# zh-cn / en / ja / es-mx / id / pt-br / ko +# BYTEDANCE_TTS_EXPLICIT_LANGUAGE=zh-cn + +# 复刻 2.0 启用字级时间戳(可选) +# BYTEDANCE_TTS_ENABLE_SUBTITLE=false + +# OpenAI兼容接口的API密钥(可选,多个用逗号分隔) OPENAI_TTS_API_KEY=your_openai_compatible_key_here -# 服务监听端口,默认8080 +# CORS 跨域白名单(逗号分隔;开发环境可设 *;空则拒绝所有跨域) +# ALLOWED_ORIGINS=https://example.com,https://app.example.com + +# 服务监听端口,默认8080 PORT=8080 diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..e7ef1d8 --- /dev/null +++ b/.gitignore @@ -0,0 +1,14 @@ +# Go build cache +.gocache/ +*.exe +*.test +*.out + +# Editor / OS +.vscode/ +.idea/ +.DS_Store +Thumbs.db + +# Logs +*.log \ No newline at end of file diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..101b7c5 --- /dev/null +++ b/Dockerfile @@ -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"] diff --git a/README.md b/README.md index bd61807..5ac04b1 100644 --- a/README.md +++ b/README.md @@ -2,215 +2,304 @@ ## 项目简介 -本项目将字节跳动火山引擎TTS(文本转语音)v3 API封装为OpenAI兼容的TTS API接口,使原本调用OpenAI TTS服务的应用可以无缝切换到火山引擎TTS服务。 +本项目将字节跳动火山引擎TTS(文本转语音)v3 API 封装为 OpenAI 兼容的 TTS API 接口,使原本调用 OpenAI TTS 服务的应用可以无缝切换到火山引擎。 ### 主要特性 -- ✅ 完全兼容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 HTTP Chunked 单向流式 API +- 支持多种音频格式:mp3 / ogg_opus / pcm / wav(wav 内部转 pcm 后本地拼头) +- 支持火山复刻 2.0 子模型(`seed-tts-2.0-standard` / `-expressive`) +- API Key 鉴权、IP 速率限制、全局并发限制 +- 内置 Prometheus 文本格式 `/metrics` 端点,零外部依赖 +- 跨平台支持(Windows / Linux / macOS) ## 快速开始 ### 前置要求 -- Go 1.19 或更高版本 -- 火山引擎账号并开通TTS服务 +- Go 1.26 或更高版本 +- 火山引擎账号并开通 TTS 服务 -### 1. 编译程序 +### 1. 编译 ```bash -go build -o tts_server tts_server.go +go build -o tts-api . ``` ### 2. 配置环境变量 -复制 `.env.example` 为 `.env` 并填入你的配置: +复制 `.env.example` 为 `.env` 并填入实际配置: ```bash cp .env.example .env ``` -编辑 `.env` 文件,填入必要的配置参数。 - -### 3. 启动服务 +### 3. 启动 ```bash # Windows -tts_server.exe +tts-api.exe # Linux/macOS -./tts_server +./tts-api ``` -服务默认监听 `8080` 端口。 +服务默认监听 `8080` 端口,可通过 `PORT` 环境变量修改。 ## 环境变量配置 ### 必需参数 -| 变量名 | 说明 | 示例 | -|--------|------|------| -| `BYTEDANCE_TTS_API_KEY` | 火山引擎新版控制台 API Key | `your_api_key_here` | -| `BYTEDANCE_TTS_RESOURCE_ID` | 资源ID,决定模型版本 | `seed-tts-1.0` | -| `BYTEDANCE_TTS_SPEAKER` | 发音人(音色)ID | `zh_female_qingxin` | +| 变量名 | 说明 | +|--------|------| +| `BYTEDANCE_TTS_API_KEY` | 火山引擎新版控制台 API Key | +| `BYTEDANCE_TTS_RESOURCE_ID` | 资源 ID,决定模型版本与计费(`seed-tts-1.0` / `seed-icl-2.0` 等) | +| `BYTEDANCE_TTS_SPEAKER` | 发音人(音色)ID,复刻音色以 `S_` 开头 | -### 可选参数 +### TTS 行为参数 | 变量名 | 说明 | 默认值 | |--------|------|--------| -| `BYTEDANCE_TTS_TIMEOUT` | 请求超时时间 | `30s` | -| `OPENAI_TTS_API_KEY` | OpenAI兼容接口的API密钥(逗号分隔支持多个) | 无 | +| `BYTEDANCE_TTS_TIMEOUT` | 单次合成超时 | `30s` | +| `BYTEDANCE_TTS_FORMAT` | 上游实际请求的音频格式(mp3 / pcm / ogg_opus);客户端要求 wav 时内部自动转 pcm + 本地拼 WAV 头 | `mp3` | +| `BYTEDANCE_TTS_SAMPLE_RATE` | 上游采样率(8000 / 16000 / 22050 / 24000 / 32000 / 44100 / 48000) | `24000` | +| `BYTEDANCE_TTS_BIT_RATE` | MP3 比特率,仅 mp3 生效 | 无 | + +### 复刻 2.0 扩展参数 + +| 变量名 | 说明 | 默认值 | +|--------|------|--------| +| `BYTEDANCE_TTS_MODEL` | 复刻 2.0 子模型(`seed-tts-2.0-standard` / `seed-tts-2.0-expressive`) | 控制台默认 | +| `BYTEDANCE_TTS_MODEL_TYPE` | 模型类型(4=ICL V2, 5=ICL V3),推荐显式指定 | 无 | +| `BYTEDANCE_TTS_EXPLICIT_LANGUAGE` | 非中英文合成时指定语种(zh-cn / en / ja / es-mx / id / pt-br / ko) | 无 | +| `BYTEDANCE_TTS_ENABLE_SUBTITLE` | 启用字级时间戳(复刻 2.0 生效) | `false` | + +### 运行时 / 服务参数 + +| 变量名 | 说明 | 默认值 | +|--------|------|--------| +| `OPENAI_TTS_API_KEY` | OpenAI 兼容接口的 API Key(逗号分隔支持多个) | 无(不鉴权) | | `PORT` | 服务监听端口 | `8080` | +| `ALLOWED_ORIGINS` | CORS 跨域白名单(逗号分隔,调试可设 `*`;空则拒绝所有跨域) | 无 | ### Resource ID 说明 | Resource ID | 模型说明 | |-------------|----------| -| `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字符版 | +| `seed-tts-1.0` | 豆包语音合成模型 1.0 字符版 | +| `seed-tts-1.0-concurr` | 豆包语音合成模型 1.0 并发版 | +| `seed-tts-2.0` | 豆包语音合成模型 2.0 字符版 | +| `seed-icl-2.0` | 声音复刻 2.0 字符版 | -**注意:** 1.0音色只能搭配 `seed-tts-1.0` Resource ID,2.0音色只能搭配 `seed-tts-2.0` Resource ID。 +> 上表为通用模型名。火山控制台实际显示的资源 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`。 + +**注意:** 复刻音色(speaker 以 `S_` 开头)必须搭配对应族的 Resource ID,否则 API 返回 resource mismatched 错误。 + +## 调试日志 + +### BYTEDANCE_TTS_DEBUG + +服务运行期日志分为**始终输出**和**调试模式才输出**两类。通过 `BYTEDANCE_TTS_DEBUG` 环境变量控制调试日志开关。 + +| 值 | 行为 | +|----|------| +| 不设置 / `false` | 仅输出错误、警告、启动摘要、成功日志(默认,生产环境推荐) | +| `true` | 额外输出适配器层调试日志 | + +```bash +# 启用调试 +BYTEDANCE_TTS_DEBUG=true ./tts-api + +# 或写入 .env +echo "BYTEDANCE_TTS_DEBUG=true" >> .env +``` + +启用后启动时会打印: + +``` +调试日志已启用 BYTEDANCE_TTS_DEBUG +``` + +### 始终输出的日志 + +启动摘要、错误警告、合成成功/失败、访问日志(Logger 中间件): + +``` +[TTS-Server] config.go:238: === 环境配置汇总 === +[TTS-Server] config.go:239: 服务端口: 8080 +... +警告: TTS 合成失败 - 路径=/v1/audio/speech 客户端=... 文本长度=50 耗时=114ms 错误=... +TTS 合成成功 - 音色=zh_female_qingxin 格式=mp3 文本=50字 音频=12345字节 分片=3 耗时=1.2s +POST /v1/audio/speech 1.2.3.4:56789 200 1.2s +``` + +### 调试模式才输出的日志(`BYTEDANCE_TTS_DEBUG=true`) + +适配器层与 CORS 拦截详情: + +``` +TTS upstream: resource_id=seed-icl-2.0 speaker=zh_female_qingxin model="seed-tts-2.0-standard" format=mp3 sample_rate=24000 speech_rate=0 additions="..." +Sentence start: sequence=0, sentence=... +Sentence end: sequence=0 +TTS 合成结束, usage: text_words=5 +volcano: 忽略未识别事件 event="xxx" sequence=1 +CORS拦截: 来源="https://..." 路径=/v1/audio/speech 方法=POST 客户端=... +``` + +> **生产建议:** 默认不开 `BYTEDANCE_TTS_DEBUG`,需要排查问题时再临时开启,避免 sentence 级别日志刷屏。 + +## CORS 跨域配置 + +跨域请求由 `ALLOWED_ORIGINS` 控制,按**完整 origin**(协议 + 域名 + 端口)精确匹配: + +- `https://app.example.com` — 精确匹配一个来源 +- `https://a.com,https://b.com` — 多个来源逗号分隔 +- `*` — 允许所有来源(**不可与凭据请求共存**) +- `app.example.com` — 缺协议头,**永远不会匹配**(强制校验 `http://` / `https://` 开头) + +**典型坑:** + +1. 客户端是 `http://` 但服务端是 `https://`:浏览器按 `http://...` 的 origin 发请求,白名单里的 `https://...` 不会匹配 → 403。**客户端必须用 `https://` 开头**。 +2. `ALLOWED_ORIGINS=*` + 客户端带 `Authorization`:浏览器按规范**直接拒绝预检**(凭据 + 通配符冲突),POST 根本发不出去。 +3. 同源请求不受 CORS 限制。 ## API 使用说明 ### OpenAI 兼容接口 -**端点:** `POST /v1/audio/speech` +**端点:** `POST /v1/audio/speech` -**请求头:** +**请求头:** - `Content-Type: application/json` -- `Authorization: Bearer <你的API密钥>`(如果配置了OPENAI_TTS_API_KEY) +- `Authorization: Bearer <你的API密钥>`(如果配置了 `OPENAI_TTS_API_KEY`) + +**请求体:** -**请求体:** ```json { "model": "tts-1", - "input": "你好,这是一个测试文本", + "input": "你好,这是一个测试文本", "voice": "alloy", - "response_format": "wav", + "response_format": "mp3", "speed": 1.0 } ``` -**参数说明:** -- `model` - 模型名称(OpenAI兼容,实际不影响) -- `input` - 要合成的文本 -- `voice` - 发音人(OpenAI兼容,实际不影响) -- `response_format` - 输出格式:仅支持 `wav` -- `speed` - 语速:0.25 ~ 4.0 +**参数说明:** +- `model` — 模型名(OpenAI 兼容,实际不影响,火山侧用 `BYTEDANCE_TTS_MODEL`) +- `input` — 要合成的文本 +- `voice` — 发音人(OpenAI 兼容,实际用 `BYTEDANCE_TTS_SPEAKER`) +- `response_format` — 输出格式:`mp3`(默认)/ `opus`(映射 ogg_opus)/ `wav` / `pcm` / `aac` / `flac`(降级到 mp3) +- `speed` — 语速,0.25 ~ 4.0(火山侧转换为 speech_rate [-50, 100]) -**示例调用:** +**格式映射:** + +| 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}' \ + -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"}' \ -o output.wav ``` -### 健康检查(含统计信息) +### 健康检查 ```bash curl http://localhost:8080/health ``` -返回包含:服务状态、请求统计、错误记录、配置检查结果 +返回服务状态、版本、运行时长、内存、配置检查结果(**不鉴权**)。 ## 限流机制 -为保护上游火山引擎API,服务实现了两层限流保护: +为保护上游火山 API,服务实现两层限流: -### 1. 全局并发限制 -- **限制**:最多同时处理 **10个** TTS请求 -- **触发**:超过10个并发请求时 -- **错误码**:`503 Service Unavailable` -- **说明**:确保不超过上游API的并发限制 +### 全局并发限制 +- 最多同时处理 **10 个** TTS 请求 +- 超过返回 `503 Service Unavailable` -### 2. IP速率限制 -- **限制**:每个IP每分钟 **100个** 请求 -- **触发**:单个IP调用过于频繁 -- **错误码**:`429 Too Many Requests` -- **说明**:防止单个客户端滥用服务 +### IP 速率限制 +- 每个 IP 每分钟 **100 个** 请求 +- 超过返回 `429 Too Many Requests` -### 触发限流时的响应 -```json -{ - "error": { - "message": "Server is busy, maximum concurrent requests reached.", - "type": "concurrency_limit_error", - "code": "max_concurrent_requests" - } -} +**触发日志(始终输出):** + +``` +警告: 已达到最大并发请求数限制,拒绝请求 - 客户端IP: 1.2.3.4 +警告: 已超过IP速率限制,拒绝请求 - 客户端IP: 1.2.3.4 ``` -### 服务器日志 -触发限流时服务器会输出中文警告日志: -- `警告: 已达到最大并发请求数限制,拒绝请求 - 客户端IP: x.x.x.x` -- `警告: 已超过IP速率限制,拒绝请求 - 客户端IP: x.x.x.x` +## 观测 / Metrics -## 支持的发音人 +服务内置 Prometheus 文本格式的 `/metrics` 端点,**不鉴权**(与 `/health` 一致),可直接被 Prometheus 抓取或浏览器查看。Go 进程内埋点,零外部依赖,实现位于 `telemetry/` 与 `metrics/` 包。 -具体发音人列表请参考火山引擎官方文档: -- 1.0音色:https://www.volcengine.com/docs/6561/97454 -- 2.0音色:https://www.volcengine.com/docs/6561/1340515 +### 主要指标 -## 常见问题 +| 指标名 | 类型 | 标签 | 说明 | +|---|---|---|---| +| `tts_request_total` | counter | status, format, speaker, model | /v1/audio/speech 请求数 | +| `tts_request_duration_seconds` | histogram | status, format | 端到端延迟 | +| `tts_upstream_total` | counter | status, format, model, speaker | 上游调用数 | +| `tts_upstream_duration_seconds` | histogram | status, format | 上游调用耗时 | +| `tts_upstream_first_byte_seconds` | histogram | format | TTFB | +| `tts_upstream_chunks_total` | counter | format | 收到的音频 chunk 数 | +| `tts_upstream_audio_bytes_total` | counter | format | 实际返回字节数 | +| `tts_upstream_errors_total` | counter | code | 上游错误(code 聚合到 transport/client/server/upstream) | +| `tts_usage_text_words_total` | counter | model | 上游计费字符数 | +| `tts_concurrency_active` | gauge | | 当前在飞请求数 | +| `tts_concurrency_rejected_total` | counter | | 并发上限拒绝数 | +| `tts_ratelimit_rejected_total` | counter | | 速率限制拒绝数 | +| `tts_auth_failed_total` | counter | | API Key 鉴权失败数 | -### 1. 如何获取鉴权信息? +### Prometheus 抓取示例 -- 登录火山引擎新版控制台 -- 进入"语音合成"服务 -- 创建应用并获取API Key - -### 2. 端口被占用怎么办? - -通过环境变量修改端口: - -```bash -# Windows -set PORT=8081 && tts_server.exe - -# Linux/macOS -PORT=8081 ./tts_server +```yaml +scrape_configs: + - job_name: tts-api + static_configs: + - targets: ['localhost:8080'] ``` -### 3. 如何配置多个API密钥? +### 仪表盘 -使用逗号分隔: +`/dashboard` 展示服务状态 + 内存 + 配置信息,并内嵌 `/metrics` 预览;Grafana 等工具可直接基于上面指标做面板。 -```bash -OPENAI_TTS_API_KEY=sk-key1,sk-key2,sk-key3 -``` +## 架构 -### 4. 查看日志 +| 包 | 职责 | +|---|---| +| `main.go` | 启动入口,信号处理 | +| `telemetry/` | Counter / Gauge / Histogram + Prometheus 文本导出(零依赖) | +| `metrics/` | TTS 业务指标注册,火山适配器埋点适配 | +| `adapter/volcano/` | 火山 v3 HTTP Chunked 客户端(client/request/response/audio/errors/synthesis) | +| `controller/` | /v1/audio/speech、/health 处理器 | +| `middleware/` | SecurityHeaders、CORS、鉴权、限流、并发、日志、客户端 IP 提取 | +| `setting/` | 单一环境变量入口 + 启动汇总 | +| `common/`、`dto/` | 常量、请求/响应类型,`common.DebugLog` 控制调试日志 | +| `router/` | 路由注册 | -服务启动后会输出详细日志,包括: -- 服务启动信息 -- 配置状态 -- 请求统计信息 -- 错误详情 +## 部署 -## 部署建议 +### Linux Systemd -### Linux Systemd 服务 - -创建 `/etc/systemd/system/tts-server.service`: +创建 `/etc/systemd/system/tts-server.service`: ```ini [Unit] @@ -222,7 +311,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 @@ -230,22 +319,76 @@ RestartSec=10 WantedBy=multi-user.target ``` -启动服务: - ```bash sudo systemctl daemon-reload sudo systemctl enable tts-server sudo systemctl start tts-server ``` -## 许可证 +### Docker -本项目采用非商业用途许可协议。您可以免费使用本软件用于非商业目的,但禁止用于任何商业活动。详细条款请参阅 [LICENSE](LICENSE) 文件。 +```bash +docker compose up -d +``` + +环境变量通过 `.env` 或 `docker-compose.yml` 传入。 + +## 常见问题 + +### 1. `code=55000000, message=resource ID is mismatched with speaker related resource` + +资源/音色不匹配。修复: + +1. 火山控制台 → 语音技术 → 你的应用 → 资源管理或音色库 +2. 用控制台在线体验/调试同一对 `BYTEDANCE_TTS_RESOURCE_ID` + 音色 +3. 控制台能合成的组合才是正确的 +4. 把控制台实际显示的资源 ID 字符串(通常是 `volc.megatts.*` 格式)填到 `BYTEDANCE_TTS_RESOURCE_ID` +5. 复刻音色(speaker 以 `S_` 开头)需确认 Resource ID 已开通且与音色同族 + +### 2. 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" --data-binary "@body.json" +``` + +JSON 用单引号包,或写到文件用 `--data-binary "@file.json"`。 + +### 3. WAV 格式音频播放异常 + +流式场景下火山 API 的 wav 格式每个 chunk 都返回完整 wav header,拼接后损坏。本项目已自动处理:选择 wav 输出时,内部用 pcm 格式请求 API,本地拼装标准 wav header。如仍有问题,改用 `mp3`。 + +### 4. 调试时如何看详细日志 + +设置 `BYTEDANCE_TTS_DEBUG=true` 后重启服务,会额外输出上游请求参数、sentence 事件、CORS 拦截等。详见上文「调试日志」一节。 + +### 5. 多 API Key 配置 + +```bash +OPENAI_TTS_API_KEY=sk-key1,sk-key2,sk-key3 +``` + +### 6. 修改端口 + +```bash +PORT=8081 ./tts-api +``` ## 技术支持 -如有问题,请检查: +如有问题,请检查: + 1. 环境变量配置是否正确 -2. 网络是否能访问火山引擎TTS服务 +2. 网络是否能访问火山引擎 TTS 服务 3. 鉴权信息是否有效 -4. Resource ID与Speaker是否匹配 +4. Resource ID 与 Speaker 是否匹配 +5. `ALLOWED_ORIGINS` 是否包含前端完整 origin(含 https://) +6. 客户端请求 URL 是否以 https:// 开头 +7. 生产环境凭据是否定期轮换 +8. 复刻音色确保 Resource ID 与音色 ID 同族 +9. 音频格式是否匹配客户端解码能力(默认 mp3 兼容性最好) + +## 许可证 + +本项目采用非商业用途许可协议。详细条款请参阅 [LICENSE](LICENSE) 文件。 diff --git a/adapter/volcano/audio.go b/adapter/volcano/audio.go new file mode 100644 index 0000000..29528c3 --- /dev/null +++ b/adapter/volcano/audio.go @@ -0,0 +1,74 @@ +package volcano + +import ( + "encoding/binary" + "fmt" +) + +// 标准 PCM WAV 头(44 字节)。 +// 文档 3.3 节:流式场景不推荐 wav(会多次返回 wav header), +// 本项目策略:上游走 pcm,本地拼一次标准头,避免拼接过个 header。 +type wavHeader struct { + // RIFF chunk descriptor + ChunkID [4]byte // "RIFF" + ChunkSize uint32 // 36 + SubChunk2Size + Format [4]byte // "WAVE" + // fmt sub-chunk + Subchunk1ID [4]byte // "fmt " + Subchunk1Size uint32 // 16 for PCM + AudioFormat uint16 // 1 = PCM + NumChannels uint16 + SampleRate uint32 + ByteRate uint32 + BlockAlign uint16 + BitsPerSample uint16 + // data sub-chunk + Subchunk2ID [4]byte // "data" + Subchunk2Size uint32 +} + +// WrapWAVHeader 把 PCM 原始字节封装成完整的 WAV 字节流。 +// sampleRate 决定 WAV 头里的采样率字段;pcm 视为 16-bit 单声道 little-endian。 +func WrapWAVHeader(pcm []byte, sampleRate int) ([]byte, error) { + if sampleRate <= 0 { + return nil, fmt.Errorf("invalid sample rate %d", sampleRate) + } + const channels uint16 = 1 + const bitsPerSample uint16 = 16 + blockAlign := channels * bitsPerSample / 8 + byteRate := uint32(sampleRate) * uint32(blockAlign) + dataSize := uint32(len(pcm)) + + hdr := wavHeader{ + ChunkID: [4]byte{'R', 'I', 'F', 'F'}, + ChunkSize: 36 + dataSize, + Format: [4]byte{'W', 'A', 'V', 'E'}, + Subchunk1ID: [4]byte{'f', 'm', 't', ' '}, + Subchunk1Size: 16, + AudioFormat: 1, + NumChannels: channels, + SampleRate: uint32(sampleRate), + ByteRate: byteRate, + BlockAlign: blockAlign, + BitsPerSample: bitsPerSample, + Subchunk2ID: [4]byte{'d', 'a', 't', 'a'}, + Subchunk2Size: dataSize, + } + + out := make([]byte, 0, 44+len(pcm)) + out = append(out, hdr.ChunkID[:]...) + out = binary.LittleEndian.AppendUint32(out, hdr.ChunkSize) + out = append(out, hdr.Format[:]...) + out = append(out, hdr.Subchunk1ID[:]...) + out = binary.LittleEndian.AppendUint32(out, hdr.Subchunk1Size) + out = binary.LittleEndian.AppendUint16(out, hdr.AudioFormat) + out = binary.LittleEndian.AppendUint16(out, hdr.NumChannels) + out = binary.LittleEndian.AppendUint32(out, hdr.SampleRate) + out = binary.LittleEndian.AppendUint32(out, hdr.ByteRate) + out = binary.LittleEndian.AppendUint16(out, hdr.BlockAlign) + out = binary.LittleEndian.AppendUint16(out, hdr.BitsPerSample) + out = append(out, hdr.Subchunk2ID[:]...) + out = binary.LittleEndian.AppendUint32(out, hdr.Subchunk2Size) + out = append(out, pcm...) + return out, nil +} diff --git a/adapter/volcano/client.go b/adapter/volcano/client.go new file mode 100644 index 0000000..ef128ae --- /dev/null +++ b/adapter/volcano/client.go @@ -0,0 +1,44 @@ +package volcano + +import ( + "bytes" + "context" + "fmt" + "net/http" + "time" +) + +// HTTPClient 持有共享的 http.Client 以便复用连接(v3 keep-alive 1 分钟)。 +type HTTPClient struct { + client *http.Client +} + +// NewHTTPClient 构造默认配置的 HTTPClient。 +func NewHTTPClient() *HTTPClient { + return &HTTPClient{ + client: &http.Client{ + Transport: &http.Transport{ + MaxIdleConns: 100, + MaxIdleConnsPerHost: 20, + IdleConnTimeout: 90 * time.Second, + TLSHandshakeTimeout: 10 * time.Second, + }, + }, + } +} + +// PostStream 发送一次流式请求,返回带上下文的 *http.Response。 +// 调用方负责关闭 resp.Body。 +func (h *HTTPClient) PostStream(ctx context.Context, url string, headers map[string]string, body []byte) (*http.Response, error) { + if ctx == nil { + ctx = context.Background() + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("build request: %w", err) + } + for k, v := range headers { + req.Header.Set(k, v) + } + return h.client.Do(req) +} diff --git a/adapter/volcano/errors.go b/adapter/volcano/errors.go new file mode 100644 index 0000000..1506d9e --- /dev/null +++ b/adapter/volcano/errors.go @@ -0,0 +1,27 @@ +package volcano + +import "fmt" + +// UpstreamError 表示火山 v3 返回的 业务错误(code != 0 且 != 20000000)或传输错误。 +// 包含上游错误码,便于 telemetry 把它作为 label。 +type UpstreamError struct { + Code int + Message string + Stage string // "request"/"stream"/"http" - 出错阶段 + Wrapped error +} + +func (e *UpstreamError) Error() string { + if e.Wrapped != nil { + return fmt.Sprintf("volcano %s: code=%d %s: %v", e.Stage, e.Code, e.Message, e.Wrapped) + } + return fmt.Sprintf("volcano %s: code=%d %s", e.Stage, e.Code, e.Message) +} + +func (e *UpstreamError) Unwrap() error { return e.Wrapped } + +// IsAuth 当上游返回认证/权限类错误时返回 true。 +func (e *UpstreamError) IsAuth() bool { + return e.Code == 45000000 || e.Code == 55000000 || + e.Code == 401 || e.Code == 403 +} diff --git a/adapter/volcano/options.go b/adapter/volcano/options.go new file mode 100644 index 0000000..444ecd2 --- /dev/null +++ b/adapter/volcano/options.go @@ -0,0 +1,63 @@ +package volcano + +// Options 是火山 v3 TTS 适配器的完整调用参数集合。 +// 由 setting 包从环境变量构造,controller 直接透传,不做 OpenAI 侧映射。 +// +// 字段顺序与文档 3.x 节一致,便于对照。 +type Options struct { + // --- 鉴权 / 路由 --- + APIKey string // X-Api-Key + ResourceID string // X-Api-Resource-Id,决定模型版本与计费,如 seed-icl-2.0 + + // --- req_params 核心字段 --- + Text string + Speaker string + Model string // 可空,仅复刻 2.0 生效;env 默认 seed-tts-2.0-standard + UID string // user.uid,默认 "uid" + + // --- audio_params --- + Format string // 上游实际请求的 format:mp3 / pcm / ogg_opus + SampleRate int // 8000/16000/22050/24000/32000/44100/48000 + BitRate int // 可选,仅 MP3 生效 + SpeechRate int // [-50, 100] + LoudnessRate int // [-50, 100] + EnableSubtitle bool // 复刻 2.0 生效,返回 TTSSubtitle + EnableTimestamp bool // 复刻 1.0 生效,内嵌字级时间戳 + + // --- additions(扩展参数,JSON 字符串承载)--- + // 文档明确 additions 在请求体里必须是 string,内容是 JSON。 + // 这里直接存结构体,序列化时由 MarshalJSON 输出为 string。 + Additions *Additions +} + +// Additions 对应文档 3.4 节的扩展参数。 +// 注意:在请求体里 additions 是 JSON 字符串,所以 MarshalJSON 序列化为 string。 +type Additions struct { + ModelType *int `json:"model_type,omitempty"` // 复刻 2.0 推荐显式指定,4=ICL V2、5=ICL V3 + ContextTexts []string `json:"context_texts,omitempty"` // 语音指令 + UseTagParser *bool `json:"use_tag_parser,omitempty"` // 复刻 2.0 expressive 启用语音标签 Cot + ExplicitLanguage string `json:"explicit_language,omitempty"` // 明确语种 + ContextLanguage string `json:"context_language,omitempty"` // 参考语种 + SilenceDuration *int `json:"silence_duration,omitempty"` // 0~30000ms + EnableLanguageDetector *bool `json:"enable_language_detector,omitempty"` // 自动识别语种 + DisableMarkdownFilter *bool `json:"disable_markdown_filter,omitempty"` // 是否解析 markdown + DisableEmojiFilter *bool `json:"disable_emoji_filter,omitempty"` // 是否过滤 emoji + MaxLengthFilterParenthesis *int `json:"max_length_to_filter_parenthesis,omitempty"` + UnsupportedCharRatio *float64 `json:"unsupported_char_ratio_thresh,omitempty"` + AIGCWatermark *bool `json:"aigc_watermark,omitempty"` + AIGCMetadata any `json:"aigc_metadata,omitempty"` + CacheConfig any `json:"cache_config,omitempty"` + PostProcess any `json:"post_process,omitempty"` +} + +// IsZero 报告 Additions 是否为空(没有任何字段设置),用于在序列化前跳过 additions。 +func (a *Additions) IsZero() bool { + if a == nil { + return true + } + return a.ModelType == nil && a.ContextTexts == nil && a.UseTagParser == nil && + a.ExplicitLanguage == "" && a.ContextLanguage == "" && a.SilenceDuration == nil && + a.EnableLanguageDetector == nil && a.DisableMarkdownFilter == nil && a.DisableEmojiFilter == nil && + a.MaxLengthFilterParenthesis == nil && a.UnsupportedCharRatio == nil && + a.AIGCWatermark == nil && a.AIGCMetadata == nil && a.CacheConfig == nil && a.PostProcess == nil +} diff --git a/adapter/volcano/request.go b/adapter/volcano/request.go new file mode 100644 index 0000000..b2fd932 --- /dev/null +++ b/adapter/volcano/request.go @@ -0,0 +1,113 @@ +package volcano + +import ( + "encoding/json" + "fmt" +) + +// requestBody 是真正发到上游 v3 端点的 JSON 顶层结构。 +type requestBody 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,omitempty"` + AudioParams ttsAudioParams `json:"audio_params"` + Additions string `json:"additions,omitempty"` // 注意:字符串 +} + +type ttsAudioParams struct { + Format string `json:"format"` + SampleRate int `json:"sample_rate"` + BitRate int `json:"bit_rate,omitempty"` + SpeechRate int `json:"speech_rate"` + LoudnessRate int `json:"loudness_rate,omitempty"` + EnableSubtitle bool `json:"enable_subtitle,omitempty"` + EnableTimestamp bool `json:"enable_timestamp,omitempty"` +} + +// buildRequest 把 Options 序列化为上游请求体 JSON。 +func buildRequest(opts Options) ([]byte, error) { + if opts.Text == "" { + return nil, fmt.Errorf("volcano: text is required") + } + if opts.Speaker == "" { + return nil, fmt.Errorf("volcano: speaker is required") + } + if opts.ResourceID == "" { + return nil, fmt.Errorf("volcano: resource id is required") + } + if opts.APIKey == "" { + return nil, fmt.Errorf("volcano: api key is required") + } + + body := requestBody{ + User: ttsUser{UID: opts.UID}, + Namespace: "UnidirectionalTTS", + ReqParams: ttsReqParams{ + Text: opts.Text, + Speaker: opts.Speaker, + Model: opts.Model, + AudioParams: ttsAudioParams{ + Format: opts.Format, + SampleRate: opts.SampleRate, + BitRate: opts.BitRate, + SpeechRate: opts.SpeechRate, + LoudnessRate: opts.LoudnessRate, + EnableSubtitle: opts.EnableSubtitle, + EnableTimestamp: opts.EnableTimestamp, + }, + }, + } + + if opts.Additions != nil && !opts.Additions.IsZero() { + // 文档明确 additions 字段为 JSON 字符串。 + raw, err := json.Marshal(opts.Additions) + if err != nil { + return nil, fmt.Errorf("marshal additions: %w", err) + } + body.ReqParams.Additions = string(raw) + } + + raw, err := json.Marshal(body) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + return raw, nil +} + +// convertSpeedToSpeechRate 把 OpenAI 风格的 speed(倍率)转换为 speech_rate(百分比)。 +// speech_rate 范围 [-50, 100],对应 0.5x ~ 2.0x。 +func convertSpeedToSpeechRate(speed float64) int { + if speed <= 0 { + speed = 1.0 + } + rate := int((speed - 1.0) * 100) + if rate < -50 { + rate = -50 + } + if rate > 100 { + rate = 100 + } + return rate +} + +// resolveUpstreamFormat 决定上游实际请求的 format。 +// - 客户端要求 wav -> 上游走 pcm,我们本地拼 header +// - 其他 -> 直接用 clientFormat +// +// sampleRate 在 wav 走 pcm 的情况下也按原样传给上游(影响 PCM 的实际采样率)。 +func resolveUpstreamFormat(clientFormat string) string { + if clientFormat == "wav" { + return "pcm" + } + return clientFormat +} diff --git a/adapter/volcano/response.go b/adapter/volcano/response.go new file mode 100644 index 0000000..a2ae1c4 --- /dev/null +++ b/adapter/volcano/response.go @@ -0,0 +1,152 @@ +package volcano + +import ( + "bufio" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "time" + + "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/dto" +) + +// ParsedStream 是一次流式响应的累计结果。 +type ParsedStream struct { + AudioData []byte + Chunks int + TextWords int + FirstChunk time.Duration // 从请求发起到收到第一个 sentence chunk 的耗时 + HasUsage bool + Subtitles []dto.SubtitleEntry +} + +// ParseStream 读取 v3 chunked NDJSON 响应,按文档 5.1 节的 event 取值分类处理。 +// +// 关键修复(对比原实现):只有 event == "sentence" 才是音频帧; +// TTSSubtitle 单独收集,不会污染音频字节流。 +func ParseStream(body io.Reader, started time.Time) (*ParsedStream, error) { + out := &ParsedStream{} + scanner := bufio.NewScanner(body) + scanner.Buffer(make([]byte, 1024*1024), 8*1024*1024) + + gotFirstChunk := false + for scanner.Scan() { + line := scanner.Bytes() + if len(line) == 0 { + continue + } + var resp dto.V3TTSResponse + if err := json.Unmarshal(line, &resp); err != nil { + if common.DebugLog { + log.Printf("volcano: 解析响应行失败: %v, line=%q", err, truncateForLog(line, 200)) + } + continue + } + + if resp.Code != 0 && resp.Code != 20000000 { + return nil, &UpstreamError{ + Code: resp.Code, + Message: resp.Message, + Stage: "stream", + } + } + + if resp.Code == 20000000 { + if resp.Usage != nil { + out.TextWords = resp.Usage.TextWords + out.HasUsage = true + if common.DebugLog { + log.Printf("TTS 合成结束, usage: text_words=%d", out.TextWords) + } + } + for scanner.Scan() { + } + break + } + + // 事件分发:显式匹配已知事件,绝不把未知事件当作音频。 + switch resp.Event { + case "TTSSentenceStart": + if common.DebugLog { + log.Printf("Sentence start: sequence=%d, sentence=%s", resp.Sequence, resp.SentenceText()) + } + case "TTSSentenceEnd": + if common.DebugLog { + log.Printf("Sentence end: sequence=%d", resp.Sequence) + } + case "TTSSubtitle": + if resp.Data != "" { + out.Subtitles = append(out.Subtitles, dto.SubtitleEntry{ + Text: resp.SentenceText(), + Sequence: resp.Sequence, + }) + } + case "sentence", "": + // HTTP 单向协议下,音频帧的 event 字段可能是空也可能是 "sentence"; + // 两种都当音频处理。 + if resp.Data == "" { + continue + } + chunk, err := base64.StdEncoding.DecodeString(resp.Data) + if err != nil { + return nil, &UpstreamError{ + Code: resp.Code, + Message: fmt.Sprintf("decode audio chunk: %v", err), + Stage: "stream", + Wrapped: err, + } + } + out.AudioData = append(out.AudioData, chunk...) + out.Chunks++ + if !gotFirstChunk { + out.FirstChunk = time.Since(started) + gotFirstChunk = true + } + default: + if common.DebugLog { + log.Printf("volcano: 忽略未识别事件 event=%q sequence=%d sentence=%s data_len=%d", resp.Event, resp.Sequence, resp.SentenceText(), len(resp.Data)) + } + } + } + + if err := scanner.Err(); err != nil { + return nil, &UpstreamError{ + Code: 0, + Message: fmt.Sprintf("read stream: %v", err), + Stage: "stream", + Wrapped: err, + } + } + + if len(out.AudioData) == 0 { + return nil, &UpstreamError{ + Code: 0, + Message: "no audio data received from TTS service", + Stage: "stream", + } + } + + return out, nil +} + +// ReadErrorBody 把非 200 响应的 body 读出来用于日志。 +func ReadErrorBody(body io.Reader) string { + const max = 2048 + buf := make([]byte, max) + n, err := io.ReadFull(body, buf) + if err != nil && !errors.Is(err, io.ErrUnexpectedEOF) && !errors.Is(err, io.EOF) { + return fmt.Sprintf("read body fail: %v", err) + } + return string(buf[:n]) +} + +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) +} diff --git a/adapter/volcano/synthesis.go b/adapter/volcano/synthesis.go new file mode 100644 index 0000000..07ccbaa --- /dev/null +++ b/adapter/volcano/synthesis.go @@ -0,0 +1,187 @@ +package volcano + +import ( + "context" + "crypto/rand" + "encoding/hex" + "fmt" + "log" + "time" + + "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/dto" +) + +// MetricsRecorder 是适配器向上报告埋点的接口。 +// 适配器本身不依赖 telemetry 包,controller 在 main 启动时把 Meter 适配成实现; +// 这样测试可以注入 mock,生产可以无侵入替换成 OTel。 +type MetricsRecorder interface { + UpstreamStarted(speaker, model, format string) + UpstreamFinished(speaker, model, format, status string, duration, ttfb time.Duration, chunks, audioBytes, errCode int) + UpstreamUsage(model string, textWords int) +} + +// nopMetrics 是 MetricsRecorder 的 no-op 默认值。 +type nopMetrics struct{} + +func (nopMetrics) UpstreamStarted(string, string, string) {} +func (nopMetrics) UpstreamFinished(string, string, string, string, time.Duration, time.Duration, int, int, int) { +} +func (nopMetrics) UpstreamUsage(string, int) {} + +// Synthesis 调用火山 v3 一次,返回组装好的结果。 +// +// 入参: +// - ctx:超时控制 +// - client:复用的 HTTPClient +// - opts:从 setting 构造的完整参数(text 字段会被 text 覆盖) +// - text:本次合成的实际文本 +// - clientFormat:客户端期望的最终格式,"wav" 内部转 pcm 后本地拼 wav 头 +// - speed:OpenAI 风格的 speed(倍率,0.5~2.0) +// - mtr:可选埋点;传 nil 等价于 nopMetrics +func Synthesis( + ctx context.Context, + client *HTTPClient, + opts Options, + text string, + clientFormat string, + speed float64, + mtr MetricsRecorder, +) (*dto.SynthesisResult, error) { + if mtr == nil { + mtr = nopMetrics{} + } + opts.Text = text + opts.SpeechRate = convertSpeedToSpeechRate(speed) + + reqID := newRequestID() + + upstreamFormat := resolveUpstreamFormat(clientFormat) + opts.Format = upstreamFormat + if upstreamFormat != "pcm" && upstreamFormat != "mp3" && upstreamFormat != "ogg_opus" { + opts.Format = "mp3" + } + + started := time.Now() + mtr.UpstreamStarted(opts.Speaker, opts.Model, opts.Format) + + body, err := buildRequest(opts) + if err != nil { + mtr.UpstreamFinished(opts.Speaker, opts.Model, opts.Format, "request_error", time.Since(started), 0, 0, 0, 0) + return nil, &UpstreamError{Code: 0, Message: err.Error(), Stage: "request", Wrapped: err} + } + + headers := map[string]string{ + "Content-Type": "application/json", + "Connection": "keep-alive", + "X-Api-Resource-Id": opts.ResourceID, + "X-Api-Request-Id": reqID, + "X-Api-Key": opts.APIKey, + "X-Control-Require-Usage-Tokens-Return": "*", + } + + if common.DebugLog { + log.Printf("TTS upstream: resource_id=%s speaker=%s model=%q format=%s sample_rate=%d speech_rate=%d additions=%q", + opts.ResourceID, opts.Speaker, opts.Model, opts.Format, opts.SampleRate, opts.SpeechRate, extractAdditionsForLog(body)) + } + + resp, err := client.PostStream(ctx, "https://openspeech.bytedance.com/api/v3/tts/unidirectional", headers, body) + if err != nil { + mtr.UpstreamFinished(opts.Speaker, opts.Model, opts.Format, "transport_error", time.Since(started), 0, 0, 0, 0) + return nil, &UpstreamError{Code: 0, Message: err.Error(), Stage: "request", Wrapped: err} + } + defer resp.Body.Close() + + if resp.StatusCode != 200 { + rawBody := ReadErrorBody(resp.Body) + mtr.UpstreamFinished(opts.Speaker, opts.Model, opts.Format, fmt.Sprintf("http_%d", resp.StatusCode), time.Since(started), 0, 0, 0, resp.StatusCode) + return nil, &UpstreamError{ + Code: resp.StatusCode, + Message: fmt.Sprintf("upstream http %d: %s", resp.StatusCode, rawBody), + Stage: "http", + } + } + + parsed, err := ParseStream(resp.Body, started) + if err != nil { + ue, _ := err.(*UpstreamError) + code := 0 + if ue != nil { + code = ue.Code + } + mtr.UpstreamFinished(opts.Speaker, opts.Model, opts.Format, "stream_error", time.Since(started), 0, 0, 0, code) + return nil, err + } + + duration := time.Since(started) + + finalData := parsed.AudioData + finalFormat := clientFormat + sampleRate := opts.SampleRate + if clientFormat == "wav" { + wav, wrapErr := WrapWAVHeader(parsed.AudioData, opts.SampleRate) + if wrapErr != nil { + mtr.UpstreamFinished(opts.Speaker, opts.Model, opts.Format, "wrap_error", duration, parsed.FirstChunk, parsed.Chunks, len(parsed.AudioData), 0) + return nil, &UpstreamError{Code: 0, Message: wrapErr.Error(), Stage: "wrap", Wrapped: wrapErr} + } + finalData = wav + } + + if parsed.HasUsage { + mtr.UpstreamUsage(opts.Model, parsed.TextWords) + } + mtr.UpstreamFinished(opts.Speaker, opts.Model, opts.Format, "ok", duration, parsed.FirstChunk, parsed.Chunks, len(finalData), 0) + + log.Printf("TTS 合成成功 - 音色=%s 格式=%s 文本=%d字 音频=%d字节 分片=%d 耗时=%v", + opts.Speaker, clientFormat, len(text), len(finalData), parsed.Chunks, duration) + + return &dto.SynthesisResult{ + AudioData: finalData, + Format: finalFormat, + SampleRate: sampleRate, + ReqID: reqID, + TextWords: parsed.TextWords, + Chunks: parsed.Chunks, + AudioBytes: len(finalData), + TTFB: parsed.FirstChunk, + Duration: duration, + }, nil +} + +// newRequestID 16 字节随机 ID(hex 编码),无外部依赖。 +func newRequestID() string { + var b [16]byte + _, _ = rand.Read(b[:]) + return hex.EncodeToString(b[:]) +} + +// extractAdditionsForLog 从已编码的请求体里取 additions 字段值,便于日志展示。 +func extractAdditionsForLog(body []byte) string { + const key = "\"additions\":\"" + idx := bytesIndex(body, key) + if idx < 0 { + return "" + } + rest := body[idx+len(key):] + end := bytesIndex(rest, "\"") + if end < 0 { + return "" + } + return string(rest[:end]) +} + +func bytesIndex(haystack []byte, needle string) int { + if len(needle) == 0 { + return 0 + } +outer: + for i := 0; i+len(needle) <= len(haystack); i++ { + for j := 0; j < len(needle); j++ { + if haystack[i+j] != needle[j] { + continue outer + } + } + return i + } + return -1 +} diff --git a/common/constants.go b/common/constants.go new file mode 100644 index 0000000..2e95ea6 --- /dev/null +++ b/common/constants.go @@ -0,0 +1,24 @@ +package common + +import "time" + +// DebugLog 控制非必要日志输出;由 setting 包在启动时通过 BYTEDANCE_TTS_DEBUG 环境变量设置。 +var DebugLog bool + +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 +) diff --git a/controller/tts.go b/controller/tts.go new file mode 100644 index 0000000..f9e869e --- /dev/null +++ b/controller/tts.go @@ -0,0 +1,259 @@ +package controller + +import ( + "context" + "encoding/json" + "fmt" + "io" + "log" + "net/http" + "runtime" + "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/metrics" + "github.com/volcano-tts/tts-api/middleware" + "github.com/volcano-tts/tts-api/setting" + "github.com/volcano-tts/tts-api/telemetry" +) + +var ( + volcanoClient *volcano.HTTPClient + adapterRec volcano.MetricsRecorder = metrics.AdapterRecorder{} +) + +func InitController() { + volcanoClient = volcano.NewHTTPClient() +} + +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) +} + +// resolveClientFormat 把 OpenAI 风格的 response_format 映射为最终输出格式; +// 不识别或未指定时回退到 setting.TTSOptions.Format。 +func resolveClientFormat(reqFmt string) string { + switch strings.ToLower(reqFmt) { + case "mp3", "wav", "opus", "pcm", "aac", "flac": + if reqFmt == "opus" { + return "ogg_opus" + } + return strings.ToLower(reqFmt) + } + if reqFmt == "" { + return setting.TTSOptions.Format + } + return setting.TTSOptions.Format +} + +// OpenaiTTSHandler 是 /v1/audio/speech 的入口。 +func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) { + start := time.Now() + + if r.Method != http.MethodPost { + log.Printf("警告: 错误的方法 - 方法=%s 期望=POST 路径=%s 客户端=%s", + r.Method, r.URL.Path, middleware.GetClientIP(r)) + metrics.RequestTotal.Inc(telemetry.Labels{"status": "method_not_allowed", "format": "", "speaker": "", "model": ""}) + http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) + return + } + + if !middleware.ValidateAPIKey(r) { + metrics.AuthFailed.Inc(telemetry.Labels{}) + 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 + } + + clientFormat := resolveClientFormat(req.ResponseFormat) + + opts := setting.TTSOptions + opts.Text = req.Input + + ctx, cancel := context.WithTimeout(r.Context(), setting.TTSTimeout) + defer cancel() + + result, err := volcano.Synthesis(ctx, volcanoClient, opts, req.Input, clientFormat, speed, adapterRec) + duration := time.Since(start) + + finalLabels := telemetry.Labels{ + "format": clientFormat, + "speaker": opts.Speaker, + "model": opts.Model, + } + if err != nil { + finalLabels["status"] = classifyStatus(err) + metrics.RequestTotal.Inc(finalLabels) + metrics.RequestDuration.Observe(duration.Seconds(), telemetry.Labels{"status": finalLabels["status"], "format": clientFormat}) + log.Printf("警告: TTS 合成失败 - 路径=%s 客户端=%s 文本长度=%d 耗时=%v 错误=%v", + r.URL.Path, middleware.GetClientIP(r), len(req.Input), duration, err) + middleware.SendJSONError(w, http.StatusInternalServerError, "TTS synthesis failed.", "server_error", "synthesis_failed") + return + } + + finalLabels["status"] = "ok" + metrics.RequestTotal.Inc(finalLabels) + metrics.RequestDuration.Observe(duration.Seconds(), telemetry.Labels{"status": "ok", "format": clientFormat}) + + w.Header().Set("Content-Type", contentTypeFor(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 classifyStatus(err error) string { + if ue, ok := err.(*volcano.UpstreamError); ok { + switch ue.Stage { + case "request": + return "request_error" + case "http": + return fmt.Sprintf("http_%d", ue.Code) + case "stream": + return "upstream_error" + case "wrap": + return "wrap_error" + } + } + return "internal_error" +} + +func contentTypeFor(format string) string { + switch strings.ToLower(format) { + case "wav": + return "audio/wav" + case "mp3": + return "audio/mpeg" + case "ogg_opus", "opus": + return "audio/ogg" + case "pcm": + return "audio/L16" + case "aac": + return "audio/aac" + case "flac": + return "audio/flac" + } + return "application/octet-stream" +} + +// HealthHandler 暴露运行期状态;无鉴权。 +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) + } + + env := setting.CheckEnvironmentVariables() + allRequired := env["all_required_vars_set"].(bool) + + status := "ok" + if !allRequired { + status = "configuration_error" + } + + resp := 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: collectMemorySnapshot(), + ConfigStatus: dto.ConfigStatusResponse{ + AllRequiredVarsSet: allRequired, + ConfigError: setting.TTSConfigErr != nil, + }, + } + json.NewEncoder(w).Encode(resp) +} + +var startTime time.Time + +func SetStartTime(t time.Time) { startTime = t } + +func collectMemorySnapshot() map[string]interface{} { + var ms runtime.MemStats + runtime.ReadMemStats(&ms) + return map[string]interface{}{ + "heap_alloc": ms.HeapAlloc, + "heap_inuse": ms.HeapInuse, + "goroutines": runtime.NumGoroutine(), + } +} diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..87e7105 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,23 @@ +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} + - OPENAI_TTS_API_KEY=${OPENAI_TTS_API_KEY:-} + - ALLOWED_ORIGINS=${ALLOWED_ORIGINS:-} + - PORT=8080 + restart: unless-stopped + healthcheck: + test: ["CMD", "wget", "-qO-", "http://localhost:8080/health"] + interval: 30s + timeout: 5s + retries: 3 + start_period: 5s diff --git a/dto/health.go b/dto/health.go new file mode 100644 index 0000000..5d844d9 --- /dev/null +++ b/dto/health.go @@ -0,0 +1,19 @@ +package dto + +// HealthResponse 是 /health 端点的 JSON 响应。 +// 数值类信息(请求统计、错误)迁移到 /metrics 端点, +// 这里只保留运行期最关键的状态。 +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"` + ConfigStatus ConfigStatusResponse `json:"config_status"` +} + +type ConfigStatusResponse struct { + AllRequiredVarsSet bool `json:"all_required_vars_set"` + ConfigError bool `json:"config_error"` +} diff --git a/dto/tts.go b/dto/tts.go new file mode 100644 index 0000000..b0420cf --- /dev/null +++ b/dto/tts.go @@ -0,0 +1,91 @@ +package dto + +import ( + "encoding/json" + "time" +) + +// OpenAITTSRequest 是 /v1/audio/speech 接收的请求体。 +// 仅 input / speed / response_format 实际影响火山侧; +// voice / model 当前保留接收但不做映射,详见 controller。 +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"` +} + +// V3TTSResponse 是火山 v3 HTTP Chunked 流式响应中每一行的 JSON 结构。 +// Sentence 字段上游有时返回字符串(TTSSentenceStart 里的句文本),有时返回对象 +// ({"phonemes":[...],"text":"...","words":[...]}),用 json.RawMessage 兼容两种形态, +// 避免任意一种上游变更都导致整行解析失败。 +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 json.RawMessage `json:"sentence,omitempty"` + IsFinal bool `json:"is_final"` + Usage *V3Usage `json:"usage,omitempty"` +} + +// SentenceText 从 Sentence 提取可读文本: +// - 字符串直接返回 +// - 对象尝试取 .text 字段 +// - 其它情况返回原始 JSON +func (r *V3TTSResponse) SentenceText() string { + if len(r.Sentence) == 0 { + return "" + } + var s string + if err := json.Unmarshal(r.Sentence, &s); err == nil { + return s + } + var obj struct { + Text string `json:"text"` + } + if err := json.Unmarshal(r.Sentence, &obj); err == nil && obj.Text != "" { + return obj.Text + } + return string(r.Sentence) +} + +// V3Usage 由 X-Control-Require-Usage-Tokens-Return 触发,包含计费字符数。 +type V3Usage struct { + TextWords int `json:"text_words"` +} + +// ByteDanceTTSConfig 是 setting 包的全局 TTS 配置,目前只承载鉴权 / URL / 超时; +// 完整的合成参数见 adapter/volcano.Options。 +type ByteDanceTTSConfig struct { + ApiKey string + ResourceId string + URL string + Timeout time.Duration +} + +// SynthesisResult 是火山适配器向 controller 返回的最终结果。 +// Format 与 AudioData 的实际编码一致;controller 据此设置响应 Content-Type。 +type SynthesisResult struct { + AudioData []byte + Format string + SampleRate int + ReqID string + TextWords int // 来自 V3Usage,无 usage 时为 0 + Chunks int // 实际收到的音频 chunk 数 + AudioBytes int // 解码后总字节数 + TTFB time.Duration // 收到首个音频 chunk 的耗时 + Duration time.Duration // 整体合成耗时 +} + +// SubtitleEntry 描述一个字级时间戳条目(当 enable_subtitle / enable_timestamp 启用时返回)。 +type SubtitleEntry struct { + Text string + StartMs int + EndMs int + Sequence int + // 原始事件可能为不同形态,这里只保留通用字段 +} diff --git a/go.mod b/go.mod index 1533b9d..6e3c122 100644 --- a/go.mod +++ b/go.mod @@ -1,8 +1,5 @@ -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 - github.com/gorilla/mux v1.8.1 -) +require github.com/gorilla/mux v1.8.1 diff --git a/go.sum b/go.sum index c9af527..7128337 100644 --- a/go.sum +++ b/go.sum @@ -1,4 +1,2 @@ -github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= -github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/gorilla/mux v1.8.1 h1:TuBL49tXwgrFYWhqrNgrUNEY92u81SPhu7sTdzQEiWY= github.com/gorilla/mux v1.8.1/go.mod h1:AKf9I4AEqPTmMytcMc0KkNouC66V3BtZ4qD5fmWSiMQ= diff --git a/health.html b/health.html new file mode 100644 index 0000000..3f31830 --- /dev/null +++ b/health.html @@ -0,0 +1,668 @@ + + + + + + TTS 服务监控 + + + + + +
+
+
+ +
+

火山 TTS 服务监控

+
{{ health.service || '' }} · {{ health.version || '' }}
+
+
+
+ + + {{ health.status === 'ok' ? '运行中' : (health.status || '加载中') }} + + + {{ lastRefresh }} +
+
+ +
{{ error }}
+ +
+
+
运行时长
+
{{ formatUptime(health.uptime) }}
+
启动于 {{ formatTime(health.start_time) }}
+
+
+
总请求数
+
{{ totalRequests }}
+
自启动以来
+
+
+
成功率
+
{{ successRate }}
+
{{ okRequests }} 成功 / {{ errRequests }} 失败
+
+
+
当前并发
+
{{ concurrencyActive }}
+
在飞请求数
+
+
+
Goroutines
+
{{ health.memory?.goroutines || 0 }}
+
Go 运行时
+
+
+
堆内存
+
{{ formatBytes(health.memory?.heap_alloc) }}
+
已分配 / 容量 {{ formatBytes(health.memory?.heap_inuse) }}
+
+
+ +
+
请求 & 流量
+
+
+
+
+
请求数
+
tts_request_total · 按 status / format / speaker 拆分
+
+
+ + + + + + + + + + +
状态格式音色次数
{{ r.status }}{{ r.format }}-{{ r.speaker || "-" }}{{ r.value }}
+
暂无数据
+
+ +
+
+
+
端到端延迟
+
tts_request_duration_seconds · 50/95/99 百分位
+
+
+ + + + + + + + + + + +
状态格式p50p95p99
{{ r.status }}{{ r.format }}-{{ r.p50 }}{{ r.p95 }}{{ r.p99 }}
+
暂无数据
+
+
+
+ +
+
上游火山 API
+
+
+
+
+
上游调用
+
tts_upstream_total · 按 status / format 拆分
+
+
+ + + + + + + + + +
状态格式次数
{{ r.status }}{{ r.format }}-{{ r.value }}
+
暂无数据
+
+ +
+
+
+
首字节耗时 (TTFB)
+
tts_upstream_first_byte_seconds · 按格式拆分
+
+
+ + + + + + + + + + +
格式p50p95p99
{{ r.format }}{{ r.p50 }}{{ r.p95 }}{{ r.p99 }}
+
暂无数据
+
+ +
+
+
+
流量统计
+
chunks & 音频字节数
+
+
+ + + + + + + + + +
格式音频分片音频字节
{{ r.format }}{{ r.chunks }}{{ r.audioBytes }}
+
暂无数据
+
+ +
+
+
+
上游错误
+
tts_upstream_errors_total · 按错误码聚合
+
+
+ + + + + + + + +
错误码次数
{{ r.code }}{{ r.value }}
+
暂无错误
+
+
+
+ +
+
限流 & 计费
+
+
+
+
+
被拒请求
+
限流 / 并发 / 鉴权失败
+
+
+ + + + + + + + + + + + + + + +
tts_ratelimit_rejected_total{{ rateLimitRejected }}
tts_concurrency_rejected_total{{ concurrencyRejected }}
tts_auth_failed_total{{ authFailed }}
+
+ +
+
+
+
计费字符
+
tts_usage_text_words_total · 按模型拆分
+
+
+ + + + + + + + +
模型字符数
{{ r.model }}{{ r.value }}
+
暂无数据
+
+
+
+
+ + + + diff --git a/main.go b/main.go new file mode 100644 index 0000000..25dc991 --- /dev/null +++ b/main.go @@ -0,0 +1,69 @@ +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/metrics" + "github.com/volcano-tts/tts-api/middleware" + "github.com/volcano-tts/tts-api/router" + "github.com/volcano-tts/tts-api/setting" +) + +func main() { + log.SetFlags(log.LstdFlags | log.Lshortfile) + log.SetPrefix("[TTS-Server] ") + + setting.InitAllConfigs() + metrics.Init() + middleware.InitRateLimiter() + setting.CheckStaticFiles() + controller.InitController() + 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("Metrics: http://localhost:%s/metrics", 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") + } +} diff --git a/metrics/metrics.go b/metrics/metrics.go new file mode 100644 index 0000000..334505d --- /dev/null +++ b/metrics/metrics.go @@ -0,0 +1,159 @@ +// Package metrics 集中声明本服务所有埋点指标,并提供 telemetry.Meter 的全局访问入口。 +// +// 设计: +// - 启动期 Init() 一次性注册所有指标;Panic 表示有重名 bug,应立即暴露。 +// - 上游适配器通过 AdapterRecorder 接入,无需直接 import telemetry。 +// - 控制器 / 中间件通过本包的全局变量直接 Inc/Observe/Set。 +package metrics + +import ( + "time" + + "github.com/volcano-tts/tts-api/telemetry" +) + +var ( + // Meter 全局 telemetry Meter。 + Meter telemetry.Meter = telemetry.NoopMeter{} + + // HTTP 请求侧 + RequestTotal *telemetry.Counter + RequestDuration *telemetry.Histogram + + // 上游 TTS 调用侧 + UpstreamTotal *telemetry.Counter + UpstreamDuration *telemetry.Histogram + UpstreamTTFB *telemetry.Histogram + UpstreamChunks *telemetry.Counter + UpstreamBytes *telemetry.Counter + UpstreamErrors *telemetry.Counter + UpstreamUsage *telemetry.Counter + + // 限流 / 并发 / 鉴权 + ConcurrencyActive *telemetry.Gauge + ConcurrencyRejected *telemetry.Counter + RateLimitRejected *telemetry.Counter + AuthFailed *telemetry.Counter +) + +// Init 初始化所有指标。在 main 启动期调用一次。 +func Init() { + m := telemetry.NewMeter() + Meter = m + + RequestTotal = m.NewCounter( + "tts_request_total", + "Total /v1/audio/speech requests, labeled by status and chosen format/speaker/model.", + "status", "format", "speaker", "model", + ) + RequestDuration = m.NewHistogram( + "tts_request_duration_seconds", + "End-to-end /v1/audio/speech latency in seconds.", + telemetry.DefaultLatencyBuckets, + "status", "format", + ) + + UpstreamTotal = m.NewCounter( + "tts_upstream_total", + "Total upstream TTS calls, labeled by status.", + "status", "format", "model", "speaker", + ) + UpstreamDuration = m.NewHistogram( + "tts_upstream_duration_seconds", + "Upstream TTS call duration in seconds.", + telemetry.DefaultLatencyBuckets, + "status", "format", + ) + UpstreamTTFB = m.NewHistogram( + "tts_upstream_first_byte_seconds", + "Time from request send to first audio chunk, in seconds.", + telemetry.DefaultLatencyBuckets, + "format", + ) + UpstreamChunks = m.NewCounter( + "tts_upstream_chunks_total", + "Total audio chunks received from upstream.", + "format", + ) + UpstreamBytes = m.NewCounter( + "tts_upstream_audio_bytes_total", + "Total audio bytes (post-wrap) returned to clients.", + "format", + ) + UpstreamErrors = m.NewCounter( + "tts_upstream_errors_total", + "Upstream TTS errors, labeled by error code family.", + "code", + ) + UpstreamUsage = m.NewCounter( + "tts_usage_text_words_total", + "Text words charged by upstream, per model.", + "model", + ) + + ConcurrencyActive = m.NewGauge( + "tts_concurrency_active", + "Current in-flight request count.", + ) + ConcurrencyRejected = m.NewCounter( + "tts_concurrency_rejected_total", + "Requests rejected due to concurrency limit.", + ) + RateLimitRejected = m.NewCounter( + "tts_ratelimit_rejected_total", + "Requests rejected due to per-IP rate limit.", + ) + AuthFailed = m.NewCounter( + "tts_auth_failed_total", + "Requests rejected due to invalid/missing API key.", + ) +} + +// AdapterRecorder 把 telemetry 指标适配为 volcano.MetricsRecorder。 +type AdapterRecorder struct{} + +// UpstreamStarted 满足 volcano.MetricsRecorder 接口。 +func (AdapterRecorder) UpstreamStarted(speaker, model, format string) { + UpstreamTotal.Inc(telemetry.Labels{"status": "started", "format": format, "model": model, "speaker": speaker}) +} + +// UpstreamFinished 满足 volcano.MetricsRecorder 接口。 +func (AdapterRecorder) UpstreamFinished(speaker, model, format, status string, duration, ttfb time.Duration, chunks, audioBytes, errCode int) { + labels := telemetry.Labels{"status": status, "format": format, "model": model, "speaker": speaker} + UpstreamTotal.Inc(labels) + UpstreamDuration.Observe(duration.Seconds(), telemetry.Labels{"status": status, "format": format}) + if ttfb > 0 { + UpstreamTTFB.Observe(ttfb.Seconds(), telemetry.Labels{"format": format}) + } + if chunks > 0 { + UpstreamChunks.Add(float64(chunks), telemetry.Labels{"format": format}) + } + if audioBytes > 0 { + UpstreamBytes.Add(float64(audioBytes), telemetry.Labels{"format": format}) + } + if errCode != 0 { + UpstreamErrors.Inc(telemetry.Labels{"code": codeLabel(errCode)}) + } +} + +// UpstreamUsage 满足 volcano.MetricsRecorder 接口。 +func (AdapterRecorder) UpstreamUsage(model string, textWords int) { + if textWords <= 0 { + return + } + UpstreamUsage.Add(float64(textWords), telemetry.Labels{"model": model}) +} + +// codeLabel 把整数错误码格式化为 label value,聚合到 4 类便于仪表盘展示。 +func codeLabel(code int) string { + switch { + case code == 0: + return "transport" + case code >= 400 && code < 500: + return "client" + case code >= 500 && code < 600: + return "server" + default: + return "upstream" + } +} diff --git a/middleware/auth.go b/middleware/auth.go new file mode 100644 index 0000000..95f6a37 --- /dev/null +++ b/middleware/auth.go @@ -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, + }, + }) +} diff --git a/middleware/cors.go b/middleware/cors.go new file mode 100644 index 0000000..79ba784 --- /dev/null +++ b/middleware/cors.go @@ -0,0 +1,101 @@ +package middleware + +import ( + "log" + "net/http" + "strings" + + "github.com/volcano-tts/tts-api/common" + "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 资源 + if common.DebugLog { + 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) + }) +} diff --git a/middleware/logger.go b/middleware/logger.go new file mode 100644 index 0000000..ea53db2 --- /dev/null +++ b/middleware/logger.go @@ -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) + }) +} diff --git a/middleware/ratelimit.go b/middleware/ratelimit.go new file mode 100644 index 0000000..3fafc5d --- /dev/null +++ b/middleware/ratelimit.go @@ -0,0 +1,150 @@ +package middleware + +import ( + "log" + "net" + "net/http" + "strings" + "sync" + "time" + + "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/metrics" +) + +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 + metrics.RateLimitRejected.Inc(nil) + 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) + } + } +} + +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 +} + +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 +} diff --git a/middleware/ratelimit_instrumented.go b/middleware/ratelimit_instrumented.go new file mode 100644 index 0000000..a9e05ef --- /dev/null +++ b/middleware/ratelimit_instrumented.go @@ -0,0 +1,58 @@ +package middleware + +// 本文件提供带 metrics 埋点的限流 / 并发中间件版本; +// 由于原 ratelimit_middleware.go 在本仓库的云盘同步下被永久占用, +// 这里用独立实现覆盖路由使用入口,旧实现保留为未引用代码。 +// +// 行为与原 ratelimit_middleware.go 完全一致,只是多了 metrics 调用。 + +import ( + "log" + "net/http" + "strings" + + "github.com/volcano-tts/tts-api/metrics" +) + +// RateLimitWithMetrics 是 middleware.RateLimit 的可埋点版本。 +func RateLimitWithMetrics(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // 仅对 /v1/ 下的业务请求限流,/health /metrics /dashboard 等监控路径不限流 + if !strings.HasPrefix(r.URL.Path, "/v1/") { + next.ServeHTTP(w, r) + return + } + 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) + }) +} + +// ConcurrencyLimitWithMetrics 是 middleware.ConcurrencyLimit 的可埋点版本。 +func ConcurrencyLimitWithMetrics(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // 仅对 /v1/ 下的业务请求统计并发和加锁,监控路径不占用并发槽位 + if !strings.HasPrefix(r.URL.Path, "/v1/") { + next.ServeHTTP(w, r) + return + } + select { + case ConcurrencySem <- struct{}{}: + metrics.ConcurrencyActive.Inc(nil) + defer func() { + <-ConcurrencySem + metrics.ConcurrencyActive.Dec(nil) + }() + next.ServeHTTP(w, r) + default: + metrics.ConcurrencyRejected.Inc(nil) + 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 + } + }) +} diff --git a/middleware/ratelimit_middleware.go b/middleware/ratelimit_middleware.go new file mode 100644 index 0000000..bbbb9f0 --- /dev/null +++ b/middleware/ratelimit_middleware.go @@ -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 + } + }) +} diff --git a/middleware/ratelimit_middleware.go.tmp b/middleware/ratelimit_middleware.go.tmp new file mode 100644 index 0000000..bbbb9f0 --- /dev/null +++ b/middleware/ratelimit_middleware.go.tmp @@ -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 + } + }) +} diff --git a/middleware/security.go b/middleware/security.go new file mode 100644 index 0000000..c3a2ff7 --- /dev/null +++ b/middleware/security.go @@ -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" || r.URL.Path == "/dashboard" || r.URL.Path == "/metrics" { + w.Header().Set("Cache-Control", "no-store") + } + + next.ServeHTTP(w, r) + }) +} diff --git a/router/router.go b/router/router.go new file mode 100644 index 0000000..e06a38d --- /dev/null +++ b/router/router.go @@ -0,0 +1,34 @@ +package router + +import ( + "net/http" + + "github.com/gorilla/mux" + "github.com/volcano-tts/tts-api/controller" + "github.com/volcano-tts/tts-api/metrics" + "github.com/volcano-tts/tts-api/middleware" +) + +func Setup() *mux.Router { + r := mux.NewRouter() + + r.Use(middleware.SecurityHeaders) + r.Use(middleware.RateLimitWithMetrics) + r.Use(middleware.ConcurrencyLimitWithMetrics) + 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") + + // /metrics 不做鉴权(对齐 /health 策略),但仍然走 RateLimit / ConcurrencyLimit。 + // Prometheus 抓取不带 Origin,因此经过 CORS 中间件时会直接 pass-through。 + r.Handle("/metrics", metrics.Meter.Handler()).Methods("GET") + + return r +} diff --git a/setting/config.go b/setting/config.go new file mode 100644 index 0000000..c48b3b3 --- /dev/null +++ b/setting/config.go @@ -0,0 +1,311 @@ +package setting + +import ( + "fmt" + "log" + "os" + "strconv" + "strings" + "time" + + "github.com/volcano-tts/tts-api/adapter/volcano" + "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/dto" +) + +// 全部环境变量读取的单一入口:其它包不允许直接 os.Getenv,只读这里的全局 Config。 + +// TTSOptions 是火山 v3 TTS 调用的完整参数集合,启动期由 InitTTSConfig 填充。 +// 业务侧(controller)直接读取并传入 volcano.Synthesis。 +var ( + TTSOptions volcano.Options + TTSConfigErr error + // TTSTimeout 单次合成请求的超时;controller 用来派生 context。 + TTSTimeout time.Duration = common.DefaultTimeout +) + +// 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 集中初始化所有配置,启动期调用一次。 +func InitAllConfigs() { + InitServerConfig() + InitAuthConfig() + InitCORSConfig() + TTSConfigErr = InitTTSConfig() +} + +func InitServerConfig() { + Server.Port = os.Getenv("PORT") + if Server.Port == "" { + Server.Port = common.DefaultPort + } +} + +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 +} + +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)) + } +} + +func normalizeOrigin(origin string) string { + origin = strings.TrimSpace(origin) + origin = strings.TrimRight(origin, "/") + return strings.ToLower(origin) +} + +// InitTTSConfig 读取火山 TTS 必填和可选配置,填充 TTSOptions 与 TTSTimeout。 +// 必填项缺失时返回 error,/v1/audio/speech 路由会拒绝请求。 +func InitTTSConfig() error { + apiKey := os.Getenv("BYTEDANCE_TTS_API_KEY") + resourceId := os.Getenv("BYTEDANCE_TTS_RESOURCE_ID") + speaker := os.Getenv("BYTEDANCE_TTS_SPEAKER") + missing := []string{} + if apiKey == "" { + missing = append(missing, "BYTEDANCE_TTS_API_KEY") + } + if resourceId == "" { + missing = append(missing, "BYTEDANCE_TTS_RESOURCE_ID") + } + if speaker == "" { + missing = append(missing, "BYTEDANCE_TTS_SPEAKER") + } + if len(missing) > 0 { + return fmt.Errorf("缺少必需的环境变量: %v", missing) + } + + model := os.Getenv("BYTEDANCE_TTS_MODEL") + format := getEnvDefault("BYTEDANCE_TTS_FORMAT", "mp3") + sampleRate := getEnvInt("BYTEDANCE_TTS_SAMPLE_RATE", 24000) + bitRate := getEnvInt("BYTEDANCE_TTS_BIT_RATE", 0) + modelType := getEnvInt("BYTEDANCE_TTS_MODEL_TYPE", 0) + explicitLanguage := os.Getenv("BYTEDANCE_TTS_EXPLICIT_LANGUAGE") + enableSubtitle := getEnvBool("BYTEDANCE_TTS_ENABLE_SUBTITLE", false) + + var adds *volcano.Additions + if modelType != 0 || explicitLanguage != "" { + adds = &volcano.Additions{} + if modelType != 0 { + v := modelType + adds.ModelType = &v + } + if explicitLanguage != "" { + adds.ExplicitLanguage = explicitLanguage + } + } + + TTSTimeout = common.DefaultTimeout + if ts := os.Getenv("BYTEDANCE_TTS_TIMEOUT"); ts != "" { + if d, err := time.ParseDuration(ts); err == nil { + TTSTimeout = d + } else { + log.Printf("无效的超时设置 %q,使用默认值 %v", ts, TTSTimeout) + } + } + + common.DebugLog = getEnvBool("BYTEDANCE_TTS_DEBUG", false) + if common.DebugLog { + log.Println("调试日志已启用 BYTEDANCE_TTS_DEBUG") + } + + TTSOptions = volcano.Options{ + APIKey: apiKey, + ResourceID: resourceId, + UID: "uid", + Speaker: speaker, + Model: model, + Format: format, + SampleRate: sampleRate, + BitRate: bitRate, + SpeechRate: 0, + LoudnessRate: 0, + EnableSubtitle: enableSubtitle, + Additions: adds, + } + return nil +} + +func getEnvDefault(name, def string) string { + if v := os.Getenv(name); v != "" { + return v + } + return def +} + +func getEnvInt(name string, def int) int { + v := os.Getenv(name) + if v == "" { + return def + } + n, err := strconv.Atoi(v) + if err != nil { + log.Printf("环境变量 %s=%q 不是合法整数,使用默认 %d", name, v, def) + return def + } + return n +} + +func getEnvBool(name string, def bool) bool { + v := os.Getenv(name) + if v == "" { + return def + } + b, err := strconv.ParseBool(v) + if err != nil { + log.Printf("环境变量 %s=%q 不是合法 bool,使用默认 %v", name, v, def) + return def + } + return b +} + +// CheckEnvironmentVariables 返回 /health 用的环境变量状态快照。 +func CheckEnvironmentVariables() map[string]interface{} { + required := map[string]bool{ + "BYTEDANCE_TTS_API_KEY": TTSOptions.APIKey != "", + "BYTEDANCE_TTS_RESOURCE_ID": TTSOptions.ResourceID != "", + "BYTEDANCE_TTS_SPEAKER": TTSOptions.Speaker != "", + } + missing := []string{} + for k, ok := range required { + if !ok { + missing = append(missing, k) + } + } + optional := map[string]bool{ + "BYTEDANCE_TTS_MODEL": TTSOptions.Model != "", + "BYTEDANCE_TTS_FORMAT": TTSOptions.Format != "mp3", + "BYTEDANCE_TTS_SAMPLE_RATE": TTSOptions.SampleRate != 24000, + "BYTEDANCE_TTS_EXPLICIT_LANGUAGE": TTSOptions.Additions != nil && TTSOptions.Additions.ExplicitLanguage != "", + "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(missing) == 0, + "missing_required_vars": missing, + "required_vars_set": required, + "optional_vars_set": optional, + } +} + +// LogStartupSummary 启动期一次性打印所有 Config 状态。 +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)) + } + + log.Printf("火山 TTS 必填项状态:") + type ttsCheck struct { + name string + value string + ok bool + } + checks := []ttsCheck{ + {"BYTEDANCE_TTS_API_KEY", maskAPIKey(TTSOptions.APIKey), TTSOptions.APIKey != ""}, + {"BYTEDANCE_TTS_RESOURCE_ID", TTSOptions.ResourceID, TTSOptions.ResourceID != ""}, + {"BYTEDANCE_TTS_SPEAKER", TTSOptions.Speaker, TTSOptions.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 整体: 初始化成功") + } +} + +func maskAPIKey(key string) string { + if key == "" { + return "" + } + if len(key) <= 8 { + return "****" + } + return key[:4] + "****" + key[len(key)-4:] +} + +// CheckStaticFiles 检查 /dashboard 路由依赖的 health.html 是否存在。 +func CheckStaticFiles() { + if _, err := os.Stat("health.html"); os.IsNotExist(err) { + log.Println("警告: health.html 不存在,/dashboard 路由将返回 404") + } +} + +// 保留 dto.ByteDanceTTSConfig 引用避免 import 警告; +// 新代码不应再使用这个类型,设置已在 TTSOptions 中。 +var _ = dto.ByteDanceTTSConfig{} diff --git a/telemetry/counter.go b/telemetry/counter.go new file mode 100644 index 0000000..1874903 --- /dev/null +++ b/telemetry/counter.go @@ -0,0 +1,90 @@ +package telemetry + +import ( + "fmt" + "io" + "sort" + "sync" + "sync/atomic" +) + +// Counter 单调递增的累计指标(整数语义,内部用 float64 位以 atomic 操作)。 +type Counter struct { + metricName string + help string + labelNames []string + + mu sync.RWMutex + values map[string]*counterChild // key = labelKey(...) +} + +type counterChild struct { + labels Labels + bits atomic.Uint64 // float64 +} + +func newCounter(name, help string, labelNames []string) *Counter { + return &Counter{ + metricName: name, + help: help, + labelNames: append([]string(nil), labelNames...), + values: make(map[string]*counterChild), + } +} + +// Inc 计数 +1。 +func (c *Counter) Inc(labels Labels) { c.Add(1, labels) } + +// Add 累加 v(v 必须 >= 0)。 +func (c *Counter) Add(v float64, labels Labels) { + if v < 0 { + return + } + child := c.getOrCreate(labels) + for { + bits := child.bits.Load() + cur := float64frombits(bits) + next := float64bits(cur + v) + if child.bits.CompareAndSwap(bits, next) { + return + } + } +} + +func (c *Counter) getOrCreate(labels Labels) *counterChild { + key := labelKey(c.labelNames, labels) + c.mu.RLock() + if child, ok := c.values[key]; ok { + c.mu.RUnlock() + return child + } + c.mu.RUnlock() + + c.mu.Lock() + defer c.mu.Unlock() + if child, ok := c.values[key]; ok { + return child + } + child := &counterChild{labels: copyLabels(labels, c.labelNames)} + c.values[key] = child + return child +} + +func (c *Counter) collect(w io.Writer) { + fmt.Fprintf(w, "# HELP %s %s\n", c.metricName, c.help) + fmt.Fprintf(w, "# TYPE %s counter\n", c.metricName) + + c.mu.RLock() + keys := make([]string, 0, len(c.values)) + for k := range c.values { + keys = append(keys, k) + } + sort.Strings(keys) + defer c.mu.RUnlock() + + for _, k := range keys { + child := c.values[k] + val := float64frombits(child.bits.Load()) + writeMetricLine(w, c.metricName, child.labels, val) + } +} diff --git a/telemetry/format.go b/telemetry/format.go new file mode 100644 index 0000000..3ba4ef7 --- /dev/null +++ b/telemetry/format.go @@ -0,0 +1,88 @@ +package telemetry + +import ( + "fmt" + "io" + "math" + "strconv" + "strings" +) + +// copyLabels 返回只包含 labelNames 中声明的 key 的副本,缺失补空串。 +// 这样序列化时输出顺序和数量固定。 +func copyLabels(labels Labels, names []string) Labels { + if len(names) == 0 { + return Labels{} + } + out := make(Labels, len(names)) + for _, n := range names { + out[n] = labels[n] + } + return out +} + +func mergeLabels(a, b Labels) Labels { + out := make(Labels, len(a)+len(b)) + for k, v := range a { + out[k] = v + } + for k, v := range b { + out[k] = v + } + return out +} + +// formatLabels 序列化为 `{k1="v1",k2="v2"}`;空集合返回空字符串。 +// value 内的 `\`, `"`, 换行会按 Prometheus 规范转义。 +func formatLabels(labels Labels) string { + if len(labels) == 0 { + return "" + } + keys := sortedKeys(labels) + var sb strings.Builder + sb.WriteByte('{') + for i, k := range keys { + if i > 0 { + sb.WriteByte(',') + } + sb.WriteString(k) + sb.WriteString(`="`) + sb.WriteString(escapeLabelValue(labels[k])) + sb.WriteByte('"') + } + sb.WriteByte('}') + return sb.String() +} + +func escapeLabelValue(v string) string { + if !strings.ContainsAny(v, "\\\"\n") { + return v + } + var sb strings.Builder + sb.Grow(len(v) + 2) + for i := 0; i < len(v); i++ { + switch v[i] { + case '\\': + sb.WriteString(`\\`) + case '"': + sb.WriteString(`\"`) + case '\n': + sb.WriteString(`\n`) + default: + sb.WriteByte(v[i]) + } + } + return sb.String() +} + +func writeMetricLine(w io.Writer, name string, labels Labels, value float64) { + fmt.Fprintf(w, "%s%s %s\n", name, formatLabels(labels), formatFloat(value)) +} + +func formatFloat(f float64) string { + return strconv.FormatFloat(f, 'g', -1, 64) +} + +// float64 bits 互转,封装到独立文件避免重复。 +func float64bits(f float64) uint64 { return math.Float64bits(f) } +func float64frombits(b uint64) float64 { return math.Float64frombits(b) } diff --git a/telemetry/gauge.go b/telemetry/gauge.go new file mode 100644 index 0000000..6c401d3 --- /dev/null +++ b/telemetry/gauge.go @@ -0,0 +1,96 @@ +package telemetry + +import ( + "fmt" + "io" + "sort" + "sync" + "sync/atomic" +) + +// Gauge 可增可减的瞬时值。 +type Gauge struct { + metricName string + help string + labelNames []string + + mu sync.RWMutex + values map[string]*gaugeChild +} + +type gaugeChild struct { + labels Labels + bits atomic.Uint64 +} + +func newGauge(name, help string, labelNames []string) *Gauge { + return &Gauge{ + metricName: name, + help: help, + labelNames: append([]string(nil), labelNames...), + values: make(map[string]*gaugeChild), + } +} + +// Set 直接设置当前值。 +func (g *Gauge) Set(v float64, labels Labels) { + child := g.getOrCreate(labels) + child.bits.Store(float64bits(v)) +} + +// Inc +1。 +func (g *Gauge) Inc(labels Labels) { g.Add(1, labels) } + +// Dec -1。 +func (g *Gauge) Dec(labels Labels) { g.Add(-1, labels) } + +// Add 累加 v(可负)。 +func (g *Gauge) Add(v float64, labels Labels) { + child := g.getOrCreate(labels) + for { + bits := child.bits.Load() + cur := float64frombits(bits) + next := float64bits(cur + v) + if child.bits.CompareAndSwap(bits, next) { + return + } + } +} + +func (g *Gauge) getOrCreate(labels Labels) *gaugeChild { + key := labelKey(g.labelNames, labels) + g.mu.RLock() + if c, ok := g.values[key]; ok { + g.mu.RUnlock() + return c + } + g.mu.RUnlock() + + g.mu.Lock() + defer g.mu.Unlock() + if c, ok := g.values[key]; ok { + return c + } + c := &gaugeChild{labels: copyLabels(labels, g.labelNames)} + g.values[key] = c + return c +} + +func (g *Gauge) collect(w io.Writer) { + fmt.Fprintf(w, "# HELP %s %s\n", g.metricName, g.help) + fmt.Fprintf(w, "# TYPE %s gauge\n", g.metricName) + + g.mu.RLock() + keys := make([]string, 0, len(g.values)) + for k := range g.values { + keys = append(keys, k) + } + sort.Strings(keys) + defer g.mu.RUnlock() + + for _, k := range keys { + child := g.values[k] + val := float64frombits(child.bits.Load()) + writeMetricLine(w, g.metricName, child.labels, val) + } +} diff --git a/telemetry/histogram.go b/telemetry/histogram.go new file mode 100644 index 0000000..6482105 --- /dev/null +++ b/telemetry/histogram.go @@ -0,0 +1,114 @@ +package telemetry + +import ( + "fmt" + "io" + "sort" + "sync" + "sync/atomic" +) + +// DefaultLatencyBuckets 适合 HTTP/TTS 场景的默认桶(秒)。 +var DefaultLatencyBuckets = []float64{0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30} + +// Histogram 累计分布型指标,记录观测值的分布。 +// +// 内部为每个 child 维护: +// - buckets[i] 累计计数(<= le_i 的观测数,不含 +Inf 桶) +// - count 全部观测计数 +// - sum 全部观测值之和 +type Histogram struct { + metricName string + help string + labelNames []string + buckets []float64 // 用户声明的上界,不含 +Inf + + mu sync.RWMutex + values map[string]*histChild +} + +type histChild struct { + labels Labels + buckets []atomic.Uint64 // 累计计数 + count atomic.Uint64 + sumBits atomic.Uint64 // float64 +} + +func newHistogram(name, help string, buckets []float64, labelNames []string) *Histogram { + bs := append([]float64(nil), buckets...) + sort.Float64s(bs) + return &Histogram{ + metricName: name, + help: help, + labelNames: append([]string(nil), labelNames...), + buckets: bs, + values: make(map[string]*histChild), + } +} + +// Observe 记录一个观测值。 +func (h *Histogram) Observe(v float64, labels Labels) { + child := h.getOrCreate(labels) + for { + bits := child.sumBits.Load() + cur := float64frombits(bits) + next := float64bits(cur + v) + if child.sumBits.CompareAndSwap(bits, next) { + break + } + } + child.count.Add(1) + for i, le := range h.buckets { + if v <= le { + child.buckets[i].Add(1) + } + } +} + +func (h *Histogram) getOrCreate(labels Labels) *histChild { + key := labelKey(h.labelNames, labels) + h.mu.RLock() + if c, ok := h.values[key]; ok { + h.mu.RUnlock() + return c + } + h.mu.RUnlock() + + h.mu.Lock() + defer h.mu.Unlock() + if c, ok := h.values[key]; ok { + return c + } + c := &histChild{ + labels: copyLabels(labels, h.labelNames), + buckets: make([]atomic.Uint64, len(h.buckets)), + } + h.values[key] = c + return c +} + +func (h *Histogram) collect(w io.Writer) { + fmt.Fprintf(w, "# HELP %s %s\n", h.metricName, h.help) + fmt.Fprintf(w, "# TYPE %s histogram\n", h.metricName) + + h.mu.RLock() + keys := make([]string, 0, len(h.values)) + for k := range h.values { + keys = append(keys, k) + } + sort.Strings(keys) + defer h.mu.RUnlock() + + for _, k := range keys { + child := h.values[k] + for i, le := range h.buckets { + merged := mergeLabels(child.labels, Labels{"le": formatFloat(le)}) + fmt.Fprintf(w, "%s_bucket%s %d\n", h.metricName, formatLabels(merged), child.buckets[i].Load()) + } + merged := mergeLabels(child.labels, Labels{"le": "+Inf"}) + fmt.Fprintf(w, "%s_bucket%s %d\n", h.metricName, formatLabels(merged), child.count.Load()) + sum := float64frombits(child.sumBits.Load()) + fmt.Fprintf(w, "%s_sum%s %s\n", h.metricName, formatLabels(child.labels), formatFloat(sum)) + fmt.Fprintf(w, "%s_count%s %d\n", h.metricName, formatLabels(child.labels), child.count.Load()) + } +} diff --git a/telemetry/labels.go b/telemetry/labels.go new file mode 100644 index 0000000..91e818f --- /dev/null +++ b/telemetry/labels.go @@ -0,0 +1,48 @@ +// Package telemetry 提供进程内可观测能力:Counter / Gauge / Histogram, +// 以及 Prometheus 文本格式导出。 +// +// 设计原则: +// - 零外部依赖,只使用标准库; +// - label key 在指标注册时锁定,运行期不可新增(避免 cardinality 爆炸); +// - 所有并发安全由实现保证,调用方无需加锁; +// - Meter 是高层入口,NoopMeter 用于测试。 +package telemetry + +import "sort" + +// Labels 是指标附加的标签集合。Value 在序列化时会按 Prometheus 规范转义。 +type Labels map[string]string + +// labelKey 计算一组标签的稳定 key,用于在内部 map 中唯一定位 child。 +// 缺失或多余的 label 一律视为空串,以保证 child 数量与 label 名集合一致。 +func labelKey(names []string, labels Labels) string { + if len(names) == 0 { + return "" + } + parts := make([]string, 0, len(names)*2) + for _, n := range names { + parts = append(parts, n, labels[n]) + } + return joinLabelParts(parts) +} + +func joinLabelParts(parts []string) string { + out := make([]byte, 0, 16*len(parts)) + for i, p := range parts { + if i > 0 { + out = append(out, 0) + } + out = append(out, p...) + } + return string(out) +} + +// sortedKeys 返回按字典序排列的 key,用于导出时输出稳定顺序。 +func sortedKeys(m map[string]string) []string { + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + sort.Strings(keys) + return keys +} diff --git a/telemetry/meter.go b/telemetry/meter.go new file mode 100644 index 0000000..12f178c --- /dev/null +++ b/telemetry/meter.go @@ -0,0 +1,60 @@ +package telemetry + +import "net/http" + +// Meter 是 telemetry 的高层入口,提供 Counter / Gauge / Histogram 的构造方法。 +// 启动时调用 NewMeter() 得到默认实现,测试时可换成 NoopMeter。 +// +// 设计:抽象成 interface 是为了在测试或禁用观测时能无侵入替换实现; +// 真正的注册逻辑全部委托给内部 *Registry。 +type Meter interface { + Handler() http.Handler + Registry() *Registry + NewCounter(name, help string, labelNames ...string) *Counter + NewGauge(name, help string, labelNames ...string) *Gauge + NewHistogram(name, help string, buckets []float64, labelNames ...string) *Histogram +} + +// RealMeter 是 Meter 的默认实现,内部维护一个 *Registry。 +type RealMeter struct { + reg *Registry +} + +// NewMeter 构造默认 Meter 实现。 +func NewMeter() Meter { + return &RealMeter{reg: newRegistry()} +} + +func (m *RealMeter) Handler() http.Handler { return m.reg.Handler() } + +// Registry 暴露给特殊用例(如测试断言),生产代码不应使用。 +func (m *RealMeter) Registry() *Registry { return m.reg } + +// NewCounter 注册并返回一个 Counter。 +// - name 指标名(Prometheus 风格,如 "tts_request_total") +// - help 帮助文本 +// - labelNames 注册时锁定的 label key 集合,运行期不可变 +func (m *RealMeter) NewCounter(name, help string, labelNames ...string) *Counter { + c := newCounter(name, help, labelNames) + if err := m.reg.register(name, c); err != nil { + // 注册重名是启动期 bug,直接 panic 让问题在启动时暴露。 + panic(err) + } + return c +} + +func (m *RealMeter) NewGauge(name, help string, labelNames ...string) *Gauge { + g := newGauge(name, help, labelNames) + if err := m.reg.register(name, g); err != nil { + panic(err) + } + return g +} + +func (m *RealMeter) NewHistogram(name, help string, buckets []float64, labelNames ...string) *Histogram { + h := newHistogram(name, help, buckets, labelNames) + if err := m.reg.register(name, h); err != nil { + panic(err) + } + return h +} diff --git a/telemetry/noop.go b/telemetry/noop.go new file mode 100644 index 0000000..c8e56f8 --- /dev/null +++ b/telemetry/noop.go @@ -0,0 +1,22 @@ +package telemetry + +import "net/http" + +// NoopMeter 是一个不采集、不输出的 Meter,用于单元测试或禁用观测的场景。 +// 返回的 Counter / Gauge / Histogram 实例不会被注册到任何 Registry, +// 它们的 Inc/Add/Observe 调用在本进程内没有可见效果(每次返回新的空实例)。 +// +// 实现 Meter 接口。 +type NoopMeter struct{} + +func (NoopMeter) NewCounter(string, string, ...string) *Counter { + return newCounter("", "", nil) +} +func (NoopMeter) NewGauge(string, string, ...string) *Gauge { + return newGauge("", "", nil) +} +func (NoopMeter) NewHistogram(string, string, []float64, ...string) *Histogram { + return newHistogram("", "", nil, nil) +} +func (NoopMeter) Handler() http.Handler { return http.NotFoundHandler() } +func (NoopMeter) Registry() *Registry { return nil } diff --git a/telemetry/registry.go b/telemetry/registry.go new file mode 100644 index 0000000..e839059 --- /dev/null +++ b/telemetry/registry.go @@ -0,0 +1,66 @@ +package telemetry + +import ( + "fmt" + "io" + "net/http" + "sort" + "sync" +) + +// collector 是 Counter / Gauge / Histogram 共同实现的内部接口。 +type collector interface { + collect(w io.Writer) +} + +// Registry 持有已注册的全部指标,提供 Prometheus 文本格式导出。 +type Registry struct { + mu sync.RWMutex + entries map[string]collector + order []string // 保留注册顺序,使输出可预测 +} + +func newRegistry() *Registry { + return &Registry{ + entries: make(map[string]collector), + } +} + +func (r *Registry) register(name string, c collector) error { + r.mu.Lock() + defer r.mu.Unlock() + if _, exists := r.entries[name]; exists { + return fmt.Errorf("metric %q already registered", name) + } + r.entries[name] = c + r.order = append(r.order, name) + return nil +} + +// Gather 把所有指标按注册顺序写入 w,文本格式遵循 Prometheus 0.0.4。 +func (r *Registry) Gather(w io.Writer) error { + r.mu.RLock() + order := append([]string(nil), r.order...) + defer r.mu.RUnlock() + for _, name := range order { + r.entries[name].collect(w) + } + return nil +} + +// Handler 返回标准 Prometheus 抓取端点。 +func (r *Registry) Handler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8") + _ = r.Gather(w) + }) +} + +// 注册顺序的辅助,用于测试断言。 +func (r *Registry) names() []string { + r.mu.RLock() + defer r.mu.RUnlock() + out := append([]string(nil), r.order...) + sort.Strings(out) + return out +} diff --git a/tts_server.go b/tts_server.go deleted file mode 100644 index d3e7f07..0000000 --- a/tts_server.go +++ /dev/null @@ -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") - } -}