diff --git a/.env.example b/.env.example index 52ce327..f0096f8 100644 --- a/.env.example +++ b/.env.example @@ -1,4 +1,4 @@ -# ByteDance TTS v3 API 配置示例 +# 字节火山引擎 TTS v3 API 配置示例 # 将此文件复制为 .env 并填入实际配置 # ========================================== @@ -8,31 +8,51 @@ # 火山引擎新版控制台获取的 API Key BYTEDANCE_TTS_API_KEY=your_api_key_here -# 资源信息ID(决定使用1.0还是2.0模型) -BYTEDANCE_TTS_RESOURCE_ID=seed-tts-1.0 +# 资源信息ID(决定使用1.0还是2.0模型) +# 复刻 2.0 音色(seed-icl-2.0 / 复刻 1.0 音色(seed-icl-1.0) +BYTEDANCE_TTS_RESOURCE_ID=seed-icl-2.0 -# 发音人(音色)ID +# 发音人(音色)ID BYTEDANCE_TTS_SPEAKER=your_speaker_id_here # ========================================== # 可选的环境变量 # ========================================== -# 请求超时时间,默认30秒 +# 单次合成超时,默认30s BYTEDANCE_TTS_TIMEOUT=30s -# 音频格式:mp3/ogg_opus/pcm/wav(默认mp3) -# 注意:流式场景下wav会多次返回header,内部自动用pcm请求再封装header +# 上游实际请求的音频格式:mp3 / pcm / ogg_opus +# 客户端要求 wav 时,内部自动转 pcm 上游 + 本地拼 WAV 头 BYTEDANCE_TTS_FORMAT=mp3 -# 音频采样率:8000/16000/22050/24000/32000/44100/48000(默认24000) +# 上游采样率:8000/16000/22050/24000/32000/44100/48000 BYTEDANCE_TTS_SAMPLE_RATE=24000 -# OpenAI兼容接口的API密钥(可选) +# 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 -# CORS 跨域白名单(逗号分隔,开发环境可设 *) +# CORS 跨域白名单(逗号分隔;开发环境可设 *;空则拒绝所有跨域) # ALLOWED_ORIGINS=https://example.com,https://app.example.com -# 服务监听端口,默认8080 +# 服务监听端口,默认8080 PORT=8080 diff --git a/README.md b/README.md index 57fa4f5..7e357b7 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,4 @@ -# 字节跳动火山引擎TTS v3 API 转 OpenAI 兼容接口 +# 字节跳动火山引擎TTS v3 API 转 OpenAI 兼容接口 ## 项目简介 @@ -307,6 +307,58 @@ POST /v1/audio/speech 1.2.3.4:56789 200 245ms POST /v1/audio/speech 1.2.3.4:56789 400 1ms ``` + +## 观测 / Metrics + +服务内置 Prometheus 文本格式的 `/metrics` 端点,**不鉴权**(与 `/health` 一致), +可直接被 Prometheus 抓取或浏览器查看。Go 进程内埋点,零外部依赖,实现位于 `telemetry/` 与 `metrics/` 包。 + +### 主要指标 + +| 指标名 | 类型 | 标签 | 说明 | +|---|---|---|---| +| `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 鉴权失败数 | + +### Prometheus 抓取示例 + +```yaml +scrape_configs: + - job_name: tts-api + static_configs: + - targets: ['localhost:8080'] +``` + +### 仪表盘 + +`/dashboard` 仍展示服务状态 + 内存 + 配置信息,并内嵌 `/metrics` 的预览; +Grafana 等工具可直接基于上面指标做面板。 + +## 架构 + +| 包 | 职责 | +|---|---| +| `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/` | 常量、请求/响应类型 | +| `router/` | 路由注册 | + ## 部署建议 ### Linux Systemd 服务 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..a244263 --- /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: "BidirectionalTTS", + 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..a75418f --- /dev/null +++ b/adapter/volcano/response.go @@ -0,0 +1,141 @@ +package volcano + +import ( + "bufio" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "time" + + "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 { + 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 + log.Printf("TTS 合成结束, usage: text_words=%d", out.TextWords) + } + for scanner.Scan() { + } + break + } + + // 事件分发:显式匹配已知事件,绝不把未知事件当作音频。 + switch resp.Event { + case "TTSSentenceStart": + log.Printf("Sentence start: sequence=%d, sentence=%s", resp.Sequence, resp.Sentence) + case "TTSSentenceEnd": + log.Printf("Sentence end: sequence=%d", resp.Sequence) + case "TTSSubtitle": + if resp.Data != "" { + out.Subtitles = append(out.Subtitles, dto.SubtitleEntry{ + Text: resp.Sentence, + Sequence: resp.Sequence, + }) + } + case "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 + } + case "": + // 传输帧,跳过 + default: + log.Printf("volcano: 忽略未识别事件 event=%q sequence=%d", resp.Event, resp.Sequence) + } + } + + 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..ecca9a4 --- /dev/null +++ b/adapter/volcano/synthesis.go @@ -0,0 +1,172 @@ +package volcano + +import ( + "context" + "crypto/rand" + "encoding/hex" + "fmt" + "log" + "time" + + "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 一次,返回组装好的结果。 +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": "*", + } + + 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) + + 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, []byte{'"'}) + 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/adapter/volcano/volcano.go b/adapter/volcano/volcano.go deleted file mode 100644 index 3a4fc30..0000000 --- a/adapter/volcano/volcano.go +++ /dev/null @@ -1,206 +0,0 @@ -package volcano - -import ( - "bufio" - "bytes" - "context" - "encoding/base64" - "encoding/json" - "fmt" - "io" - "log" - "net/http" - "time" - - "github.com/google/uuid" - "github.com/volcano-tts/tts-api/dto" -) - -// 火山引擎 TTS v3 HTTP 单向流式 API 客户端。 -// 官方文档:https://www.volcengine.com/docs/6561/1598757 -// 官方 Go 示例:请求体仅含 req_params;严格按单文件参考实现的请求结构。 - -type HTTPClient struct { - client *http.Client -} - -func NewHTTPClient() *HTTPClient { - return &HTTPClient{ - client: &http.Client{ - Transport: &http.Transport{ - MaxIdleConns: 100, - MaxIdleConnsPerHost: 20, - IdleConnTimeout: 90 * time.Second, - TLSHandshakeTimeout: 10 * time.Second, - }, - }, - } -} - -func (h *HTTPClient) PostStream(url string, headers map[string]string, body []byte, timeout time.Duration) (*http.Response, error) { - req, err := http.NewRequest(http.MethodPost, url, bytes.NewBuffer(body)) - if err != nil { - return nil, err - } - for key, value := range headers { - req.Header.Set(key, value) - } - - ctx, cancel := context.WithTimeout(context.Background(), timeout) - defer cancel() - - req = req.WithContext(ctx) - return h.client.Do(req) -} - -// --- 请求体结构(对应火山 v3 API 请求 JSON) --- - -type ttsRequest struct { - User ttsUser `json:"user"` - Namespace string `json:"namespace"` - ReqParams ttsReqParams `json:"req_params"` -} - -type ttsUser struct { - UID string `json:"uid"` -} - -type ttsReqParams struct { - Text string `json:"text"` - Speaker string `json:"speaker"` - AudioParams ttsAudioParams `json:"audio_params"` -} - -type ttsAudioParams struct { - Format string `json:"format"` - SampleRate int `json:"sample_rate"` - SpeechRate int `json:"speech_rate"` -} - -// convertSpeedToSpeechRate 把 OpenAI 风格的 speed(倍率)转成火山 v3 的 speech_rate(百分比)。 -// 文档规定 speech_rate 范围 [-50, 100],对应 0.5x ~ 2.0x 倍速。 -// 输入超出范围会被截断到边界值。 -func convertSpeedToSpeechRate(speed float64) int { - rate := int((speed - 1.0) * 100) - if rate < -50 { - rate = -50 - } - if rate > 100 { - rate = 100 - } - return rate -} -func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text string, speed float64) (*dto.SynthesisResult, error) { - reqID := uuid.NewString() - speechRate := convertSpeedToSpeechRate(speed) - - // 构造请求体:严格按 v3 API 文档 JSON 结构 - req := ttsRequest{ - User: ttsUser{UID: "uid"}, - Namespace: "BidirectionalTTS", - ReqParams: ttsReqParams{ - Text: text, - Speaker: config.Speaker, - AudioParams: ttsAudioParams{ - Format: "wav", - SampleRate: 24000, - SpeechRate: speechRate, - }, - }, - } - - body, err := json.Marshal(req) - if err != nil { - return nil, fmt.Errorf("marshal TTS request: %w", err) - } - // 诊断日志:记录实际发到上游的请求体(去 model/speaker/resource 关键字段) - log.Printf("TTS upstream request: X-Api-Resource-Id=%s speaker=%s namespace=BidirectionalTTS body=%s", - config.ResourceId, config.Speaker, string(body)) - - // 鉴权 header 按 v3 新版控制台方式(Connection 由 Go http 默认 keep-alive) - headers := map[string]string{ - "Content-Type": "application/json", - "Connection": "keep-alive", - "X-Api-Resource-Id": config.ResourceId, - "X-Api-Request-Id": reqID, - "X-Api-Key": config.ApiKey, - } - - resp, err := httpClient.PostStream(config.URL, headers, body, config.Timeout) - if err != nil { - return nil, fmt.Errorf("send TTS request: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - respBody, readErr := io.ReadAll(resp.Body) - if readErr != nil { - log.Printf("TTS service error: status=%d, read body fail: %v", resp.StatusCode, readErr) - } else { - log.Printf("TTS service error: status=%d, body=%s", resp.StatusCode, string(respBody)) - } - return nil, fmt.Errorf("TTS service error: status %d", resp.StatusCode) - } - - var audioData []byte - scanner := bufio.NewScanner(resp.Body) - // 初始 1MB / 最大 8MB,与示例同量级,留足 TTS 长文本 room - scanner.Buffer(make([]byte, 1024*1024), 8*1024*1024) - - for scanner.Scan() { - line := scanner.Bytes() - if len(line) == 0 { - continue - } - - var v3Resp dto.V3TTSResponse - if err := json.Unmarshal(line, &v3Resp); err != nil { - log.Printf("unmarshal chunk fail: %v, line: %s", err, string(line)) - continue - } - - // code=20000000 表示合成结束 - if v3Resp.Code == 20000000 { - if v3Resp.Usage != nil { - log.Printf("TTS synthesis completed, usage: %+v", v3Resp.Usage) - } - // 跳过后续可能的空行 - for scanner.Scan() { - } - break - } - - // 非零 code 为错误 - if v3Resp.Code != 0 { - log.Printf("TTS service error: code=%d, message=%s, event=%s", v3Resp.Code, v3Resp.Message, v3Resp.Event) - return nil, fmt.Errorf("TTS service error: %s", v3Resp.Message) - } - - // 根据 event 字段分类处理 - switch v3Resp.Event { - case "TTSSentenceStart": - log.Printf("Sentence start: sequence=%d, sentence=%s", v3Resp.Sequence, v3Resp.Sentence) - case "TTSSentenceEnd": - log.Printf("Sentence end: sequence=%d", v3Resp.Sequence) - default: - // 音频数据 chunk:data 字段为 base64 编码的音频片段 - if v3Resp.Data != "" { - chunk, err := base64.StdEncoding.DecodeString(v3Resp.Data) - if err != nil { - return nil, fmt.Errorf("decode audio chunk: %w", err) - } - audioData = append(audioData, chunk...) - } - } - } - - if err := scanner.Err(); err != nil { - return nil, fmt.Errorf("read TTS stream: %w", err) - } - - if len(audioData) == 0 { - return nil, fmt.Errorf("no audio data received from TTS service") - } - - return &dto.SynthesisResult{AudioData: audioData, ReqID: reqID, Format: "wav"}, nil -} diff --git a/controller/tts.go b/controller/tts.go index 66e4c54..49a409f 100644 --- a/controller/tts.go +++ b/controller/tts.go @@ -1,29 +1,34 @@ -package controller +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/service" "github.com/volcano-tts/tts-api/setting" + "github.com/volcano-tts/tts-api/telemetry" ) -var volcanoClient *volcano.HTTPClient +var ( + volcanoClient *volcano.HTTPClient + adapterRec volcano.MetricsRecorder = metrics.AdapterRecorder{} +) func InitController() { volcanoClient = volcano.NewHTTPClient() } -// truncateForLog 用于在日志中安全地展示请求内容(截断避免日志爆炸、控制不可打印字符) func truncateForLog(b []byte, max int) string { if len(b) > max { return string(b[:max]) + fmt.Sprintf("...(truncated, total %d bytes)", len(b)) @@ -31,15 +36,36 @@ func truncateForLog(b []byte, max int) string { 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") @@ -47,7 +73,7 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) { } if setting.TTSConfigErr != nil { - log.Printf("警告: TTS配置未就绪,拒绝请求 - 错误=%v 路径=%s 客户端=%s", + 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 @@ -96,7 +122,6 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) { 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) @@ -115,27 +140,78 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) { speed = common.MaxSpeed } - ttsStart := time.Now() - result, err := volcano.Synthesis(&setting.TTSConfig, volcanoClient, req.Input, speed) - duration := time.Since(ttsStart) + 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 { - service.GlobalStats.AddRequest(false, duration, err.Error()) + 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) http.Error(w, "TTS synthesis failed", http.StatusInternalServerError) return } - service.GlobalStats.AddRequest(true, duration, "") + finalLabels["status"] = "ok" + metrics.RequestTotal.Inc(finalLabels) + metrics.RequestDuration.Observe(duration.Seconds(), telemetry.Labels{"status": "ok", "format": clientFormat}) - w.Header().Set("Content-Type", "audio/wav") + 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") @@ -145,55 +221,39 @@ func HealthHandler(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) } - totalRequests, successfulRequests, failedRequests, totalResponseTime, recentResponseTimes, lastErrors := service.GlobalStats.GetSnapshot() - - var errorRate float64 - if totalRequests > 0 { - errorRate = float64(failedRequests) / float64(totalRequests) * 100 - } - - var avgResponseTime float64 - if totalRequests > 0 { - avgResponseTime = totalResponseTime.Seconds() * 1000 / float64(totalRequests) - } - - envCheckStatus := setting.CheckEnvironmentVariables() - allEnvVarsSet := envCheckStatus["all_required_vars_set"].(bool) + env := setting.CheckEnvironmentVariables() + allRequired := env["all_required_vars_set"].(bool) status := "ok" - if !allEnvVarsSet { + if !allRequired { status = "configuration_error" } - response := dto.HealthResponse{ + 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: service.GetMemoryInfo(), - APIStats: dto.APIStatsResponse{ - TotalRequests: int(totalRequests), - SuccessfulRequests: successfulRequests, - FailedRequests: failedRequests, - ErrorRatePercent: fmt.Sprintf("%.2f", errorRate), - AvgResponseTimeMs: fmt.Sprintf("%.2f", avgResponseTime), - RecentResponseTimesMs: recentResponseTimes, - }, - Errors: dto.ErrorResponse{ - RecentErrorsCount: len(lastErrors), - }, + Memory: collectMemorySnapshot(), ConfigStatus: dto.ConfigStatusResponse{ - AllRequiredVarsSet: allEnvVarsSet, + AllRequiredVarsSet: allRequired, ConfigError: setting.TTSConfigErr != nil, }, } - - json.NewEncoder(w).Encode(response) + json.NewEncoder(w).Encode(resp) } var startTime time.Time -func SetStartTime(t time.Time) { - startTime = t +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/dto/health.go b/dto/health.go index c1ba13b..c479603 100644 --- a/dto/health.go +++ b/dto/health.go @@ -1,5 +1,8 @@ -package dto +package dto +// HealthResponse 是 /health 端点的 JSON 响应。 +// 数值类信息(请求统计、错误)迁移到 /metrics 端点, +// 这里只保留运行期最关键的状态。 type HealthResponse struct { Status string `json:"status"` Service string `json:"service"` @@ -7,24 +10,9 @@ type HealthResponse struct { Uptime string `json:"uptime"` StartTime string `json:"start_time"` Memory map[string]interface{} `json:"memory"` - APIStats APIStatsResponse `json:"api_stats"` - Errors ErrorResponse `json:"errors"` ConfigStatus ConfigStatusResponse `json:"config_status"` } -type APIStatsResponse struct { - TotalRequests int `json:"total_requests"` - SuccessfulRequests int64 `json:"successful_requests"` - FailedRequests int64 `json:"failed_requests"` - ErrorRatePercent string `json:"error_rate_percent"` - AvgResponseTimeMs string `json:"avg_response_time_ms"` - RecentResponseTimesMs []float64 `json:"recent_response_times_ms"` -} - -type ErrorResponse struct { - RecentErrorsCount int `json:"recent_errors_count"` -} - type ConfigStatusResponse struct { AllRequiredVarsSet bool `json:"all_required_vars_set"` ConfigError bool `json:"config_error"` diff --git a/dto/tts.go b/dto/tts.go index fedaff8..402eb1f 100644 --- a/dto/tts.go +++ b/dto/tts.go @@ -1,7 +1,10 @@ -package dto +package dto import "time" +// OpenAITTSRequest 是 /v1/audio/speech 接收的请求体。 +// 仅 input / speed / response_format 实际影响火山侧; +// voice / model 当前保留接收但不做映射,详见 controller。 type OpenAITTSRequest struct { Model string `json:"model"` Input string `json:"input"` @@ -10,6 +13,7 @@ type OpenAITTSRequest struct { Speed float64 `json:"speed,omitempty"` } +// V3TTSResponse 是火山 v3 HTTP Chunked 流式响应中每一行的 JSON 结构。 type V3TTSResponse struct { ReqID string `json:"reqid"` Code int `json:"code"` @@ -22,20 +26,39 @@ type V3TTSResponse struct { Usage *V3Usage `json:"usage,omitempty"` } +// 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 - Speaker string URL string Timeout time.Duration } +// SynthesisResult 是火山适配器向 controller 返回的最终结果。 +// Format 与 AudioData 的实际编码一致;controller 据此设置响应 Content-Type。 type SynthesisResult struct { - AudioData []byte - ReqID string - Format string // 实际输出格式,用于设置 Content-Type + 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/health.html b/health.html index ff92635..2a28df5 100644 --- a/health.html +++ b/health.html @@ -1,4 +1,4 @@ - + @@ -7,241 +7,73 @@ @@ -250,14 +82,15 @@

TTS 服务监控

{{ healthData.service || 'Loading...' }} - {{ healthData.version || '' }}
- +
+ + 查看 Prometheus 指标 +
-
- {{ error }} -
+
{{ error }}
@@ -271,67 +104,25 @@
-
- - 请求统计 -
-
-
-
{{ healthData.api_stats?.total_requests || 0 }}
-
总请求数
-
-
-
{{ healthData.api_stats?.successful_requests || 0 }}
-
成功
-
-
-
{{ healthData.api_stats?.failed_requests || 0 }}
-
失败
-
-
-
- -
-
- - 错误率 -
-
- {{ healthData.api_stats?.error_rate_percent || '0' }}% -
-
平均响应: {{ healthData.api_stats?.avg_response_time_ms || '0' }} ms
-
-
- -
-
-
- - 配置状态 -
+
配置状态
环境变量 {{ healthData.config_status?.all_required_vars_set ? '已配置' : '未配置' }}
-
- 配置状态 - - {{ healthData.config_status?.config_error ? '异常' : '正常' }} - -
启动时间 {{ healthData.start_time || '-' }}
+
+ 详细指标 + /metrics +
-
- - 内存使用 -
+
内存使用
{{ formatBytes(healthData.memory?.heap_alloc) }}
@@ -347,51 +138,25 @@
- -
-
- - 响应时间趋势 -
-
-
-
-
-
-
- - 错误记录 ({{ recentErrorsCount }}) -
-
- 暂无错误记录 -
-
-
-
检测到 {{ recentErrorsCount }} 条最近错误,详细信息请查看服务器日志
-
-
+
指标预览 (Prometheus 文本格式)
+
-
- 加载中... -
+
加载中...
- \ No newline at end of file + diff --git a/main.go b/main.go index aa8a979..fd7d392 100644 --- a/main.go +++ b/main.go @@ -1,4 +1,4 @@ -package main +package main import ( "context" @@ -10,9 +10,9 @@ import ( "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/service" "github.com/volcano-tts/tts-api/setting" ) @@ -20,17 +20,11 @@ func main() { log.SetFlags(log.LstdFlags | log.Lshortfile) log.SetPrefix("[TTS-Server] ") - // 所有环境变量读取在 setting 包内集中完成,业务模块只读全局 Config。 setting.InitAllConfigs() - - // 兼容旧调用顺序:rate limiter / 静态文件 / stats / controller 的初始化保持独立。 + metrics.Init() middleware.InitRateLimiter() setting.CheckStaticFiles() - service.InitStats() controller.InitController() - - // 启动期一次性打印所有 Config 状态,便于运维核对。 - // (必填项缺失的明确警告由 LogStartupSummary 自身负责,避免重复打印。) setting.LogStartupSummary() controller.SetStartTime(time.Now()) @@ -53,6 +47,7 @@ func main() { 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 { diff --git a/metrics/metrics.go b/metrics/metrics.go new file mode 100644 index 0000000..8041eed --- /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/ratelimit.go b/middleware/ratelimit.go index a3ba367..7baee32 100644 --- a/middleware/ratelimit.go +++ b/middleware/ratelimit.go @@ -1,4 +1,4 @@ -package middleware +package middleware import ( "log" @@ -9,6 +9,7 @@ import ( "time" "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/metrics" ) type RateLimiter struct { @@ -55,6 +56,7 @@ func (rl *RateLimiter) Allow(key string) bool { if len(valid) >= rl.limit { rl.requests[key] = valid + metrics.RateLimitRejected.Inc(nil) return false } @@ -90,7 +92,6 @@ func (rl *RateLimiter) cleanup() { } } -// 私有网络 CIDR 范围:仅在直连来源属于这些范围时才信任代理头 var privateCIDRs []*net.IPNet func init() { @@ -125,9 +126,6 @@ func isPrivateIP(ipStr string) bool { return false } -// GetClientIP 提取客户端真实 IP。 -// 仅当直连来源为私有网络(本地代理、Docker 网桥等)时才信任 X-Forwarded-For / X-Real-IP, -// 防止公网直连场景下攻击者伪造代理头绕过速率限制。 func GetClientIP(r *http.Request) string { directIP, _, err := net.SplitHostPort(r.RemoteAddr) if err != nil { diff --git a/middleware/ratelimit_instrumented.go b/middleware/ratelimit_instrumented.go new file mode 100644 index 0000000..3cd7e8b --- /dev/null +++ b/middleware/ratelimit_instrumented.go @@ -0,0 +1,47 @@ +package middleware + +// 本文件提供带 metrics 埋点的限流 / 并发中间件版本; +// 由于原 ratelimit_middleware.go 在本仓库的云盘同步下被永久占用, +// 这里用独立实现覆盖路由使用入口,旧实现保留为未引用代码。 +// +// 行为与原 ratelimit_middleware.go 完全一致,只是多了 metrics 调用。 + +import ( + "log" + "net/http" + + "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) { + 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) { + 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.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/router/router.go b/router/router.go index 4635c0f..4432fa4 100644 --- a/router/router.go +++ b/router/router.go @@ -1,10 +1,11 @@ -package router +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" ) @@ -12,8 +13,8 @@ func Setup() *mux.Router { r := mux.NewRouter() r.Use(middleware.SecurityHeaders) - r.Use(middleware.RateLimit) - r.Use(middleware.ConcurrencyLimit) + r.Use(middleware.RateLimitWithMetrics) + r.Use(middleware.ConcurrencyLimitWithMetrics) r.Use(middleware.Logger) r.HandleFunc("/v1/audio/speech", controller.OpenaiTTSHandler).Methods("POST", "OPTIONS") @@ -25,5 +26,9 @@ func Setup() *mux.Router { 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/service/stats.go b/service/stats.go deleted file mode 100644 index 1fce766..0000000 --- a/service/stats.go +++ /dev/null @@ -1,124 +0,0 @@ -package service - -import ( - "runtime" - "strings" - "sync" - "time" - - "github.com/volcano-tts/tts-api/common" -) - -type Stats struct { - totalRequests int64 - successfulRequests int64 - failedRequests int64 - totalResponseTime time.Duration - recentResponseTimes []float64 - responseTimesIndex int - responseTimesCount int - lastErrors []string - errorsIndex int - errorsCount int - mutex sync.RWMutex -} - -var GlobalStats *Stats - -func InitStats() { - GlobalStats = &Stats{ - recentResponseTimes: make([]float64, common.MaxResponseTimes), - lastErrors: make([]string, common.MaxErrors), - } -} - -func (s *Stats) AddRequest(success bool, responseTime time.Duration, errMsg string) { - s.mutex.Lock() - defer s.mutex.Unlock() - - s.totalRequests++ - s.totalResponseTime += responseTime - - s.recentResponseTimes[s.responseTimesIndex] = responseTime.Seconds() * 1000 - s.responseTimesIndex = (s.responseTimesIndex + 1) % common.MaxResponseTimes - if s.responseTimesCount < common.MaxResponseTimes { - s.responseTimesCount++ - } - - if success { - s.successfulRequests++ - } else { - s.failedRequests++ - if errMsg != "" { - now := time.Now().Format(time.RFC3339) - - // 去重:如果最近一条错误的消息内容相同,仅更新时间戳 - if s.errorsCount > 0 { - lastIdx := (s.errorsIndex - 1 + common.MaxErrors) % common.MaxErrors - lastEntry := s.lastErrors[lastIdx] - if sepIdx := strings.Index(lastEntry, ": "); sepIdx != -1 { - if lastEntry[sepIdx+2:] == errMsg { - s.lastErrors[lastIdx] = now + ": " + errMsg - return - } - } - } - - s.lastErrors[s.errorsIndex] = now + ": " + errMsg - s.errorsIndex = (s.errorsIndex + 1) % common.MaxErrors - if s.errorsCount < common.MaxErrors { - s.errorsCount++ - } - } - } -} - -func (s *Stats) GetSnapshot() (totalRequests int64, successfulRequests int64, failedRequests int64, - totalResponseTime time.Duration, recentResponseTimes []float64, lastErrors []string) { - s.mutex.RLock() - defer s.mutex.RUnlock() - - totalRequests = s.totalRequests - successfulRequests = s.successfulRequests - failedRequests = s.failedRequests - totalResponseTime = s.totalResponseTime - - // 按时间顺序(从旧到新)遍历响应时间环形缓冲区 - recentResponseTimes = make([]float64, 0, s.responseTimesCount) - if s.responseTimesCount > 0 { - start := 0 - if s.responseTimesCount == common.MaxResponseTimes { - start = s.responseTimesIndex - } - for i := 0; i < s.responseTimesCount; i++ { - idx := (start + i) % common.MaxResponseTimes - recentResponseTimes = append(recentResponseTimes, s.recentResponseTimes[idx]) - } - } - - // 按时间顺序(从旧到新)遍历错误环形缓冲区 - lastErrors = make([]string, 0, s.errorsCount) - if s.errorsCount > 0 { - start := 0 - if s.errorsCount == common.MaxErrors { - start = s.errorsIndex - } - for i := 0; i < s.errorsCount; i++ { - idx := (start + i) % common.MaxErrors - lastErrors = append(lastErrors, s.lastErrors[idx]) - } - } - - return -} - -func GetMemoryInfo() map[string]interface{} { - var m runtime.MemStats - runtime.ReadMemStats(&m) - return map[string]interface{}{ - "total_alloc": m.TotalAlloc, - "heap_alloc": m.HeapAlloc, - "heap_inuse": m.HeapInuse, - "goroutines": runtime.NumGoroutine(), - } -} diff --git a/setting/config.go b/setting/config.go index bfc405d..b6dd083 100644 --- a/setting/config.go +++ b/setting/config.go @@ -1,22 +1,27 @@ -package setting +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。 -// TTSConfig 上游火山 TTS 配置(由 InitTTSConfig 填充)。 +// TTSOptions 是火山 v3 TTS 调用的完整参数集合,启动期由 InitTTSConfig 填充。 +// 业务侧(controller)直接读取并传入 volcano.Synthesis。 var ( - TTSConfig dto.ByteDanceTTSConfig + TTSOptions volcano.Options TTSConfigErr error + // TTSTimeout 单次合成请求的超时;controller 用来派生 context。 + TTSTimeout time.Duration = common.DefaultTimeout ) // AuthConfig OpenAI 兼容接口的客户端 API Key 鉴权配置。 @@ -42,8 +47,6 @@ type ServerConfig struct { var Server ServerConfig // InitAllConfigs 集中初始化所有配置,启动期调用一次。 -// 返回 TTSConfigErr(火山 TTS 必填项缺失时为非 nil);其它 Config 缺失时不返回 error, -// 各自有合理兜底(Auth 放行 / CORS 拒绝跨域 / Server 默认 8080)。 func InitAllConfigs() { InitServerConfig() InitAuthConfig() @@ -51,7 +54,6 @@ func InitAllConfigs() { TTSConfigErr = InitTTSConfig() } -// InitServerConfig 读取 PORT,缺省 common.DefaultPort。 func InitServerConfig() { Server.Port = os.Getenv("PORT") if Server.Port == "" { @@ -59,8 +61,6 @@ func InitServerConfig() { } } -// InitAuthConfig 读取 OPENAI_TTS_API_KEY,支持逗号分隔多个 key。 -// 留空时 Auth.APIKeys 为空,ValidateAPIKey 会放行所有请求。 func InitAuthConfig() { raw := os.Getenv("OPENAI_TTS_API_KEY") if raw == "" { @@ -78,8 +78,6 @@ func InitAuthConfig() { Auth.APIKeys = keys } -// InitCORSConfig 读取 ALLOWED_ORIGINS,按逗号分隔;支持 * 通配(AllowAll=true)。 -// 留空时 CORS.Origins 为空,跨域请求会被拒绝。 func InitCORSConfig() { raw := os.Getenv("ALLOWED_ORIGINS") CORS.Origins = nil @@ -100,58 +98,142 @@ func InitCORSConfig() { } } -// normalizeOrigin 复制自原 middleware/cors.go:小写 + 去尾斜杠。 func normalizeOrigin(origin string) string { origin = strings.TrimSpace(origin) origin = strings.TrimRight(origin, "/") return strings.ToLower(origin) } -// InitTTSConfig 读取火山 TTS 必填和可选配置,填充 TTSConfig。 -// 必填项缺失时返回 error,服务可继续运行但 TTS 功能不可用。 +// 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") - missingVars := []string{} + missing := []string{} if apiKey == "" { - missingVars = append(missingVars, "BYTEDANCE_TTS_API_KEY") + missing = append(missing, "BYTEDANCE_TTS_API_KEY") } if resourceId == "" { - missingVars = append(missingVars, "BYTEDANCE_TTS_RESOURCE_ID") + missing = append(missing, "BYTEDANCE_TTS_RESOURCE_ID") } if speaker == "" { - missingVars = append(missingVars, "BYTEDANCE_TTS_SPEAKER") + missing = append(missing, "BYTEDANCE_TTS_SPEAKER") + } + if len(missing) > 0 { + return fmt.Errorf("缺少必需的环境变量: %v", missing) } - if len(missingVars) > 0 { - return fmt.Errorf("缺少必需的环境变量: %v", missingVars) - } + 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) - url := "https://openspeech.bytedance.com/api/v3/tts/unidirectional" - - timeout := common.DefaultTimeout - if timeoutStr := os.Getenv("BYTEDANCE_TTS_TIMEOUT"); timeoutStr != "" { - if parsedTimeout, err := time.ParseDuration(timeoutStr); err == nil { - timeout = parsedTimeout - } else { - log.Printf("无效的超时设置 '%s',使用默认值: %v", timeoutStr, timeout) + var adds *volcano.Additions + if modelType != 0 || explicitLanguage != "" { + adds = &volcano.Additions{} + if modelType != 0 { + v := modelType + adds.ModelType = &v + } + if explicitLanguage != "" { + adds.ExplicitLanguage = explicitLanguage } } - TTSConfig = dto.ByteDanceTTSConfig{ - ApiKey: apiKey, - ResourceId: resourceId, - Speaker: speaker, - URL: url, - Timeout: timeout, + 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) + } + } + + 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 } -// LogStartupSummary 在启动期打印所有 Config 的最终状态。 -// 调用时机:InitAllConfigs 之后,ListenAndServe 之前。 -// 必填项逐项输出,失败分支明确告知"v1/audio/speech 路由将 500"。 +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) @@ -163,14 +245,13 @@ func LogStartupSummary() { } if CORS.AllowAll { - log.Printf("ALLOWED_ORIGINS: *(允许所有跨域,不可与凭据共用)") + log.Printf("ALLOWED_ORIGINS: *(允许所有跨域;不可与鉴权共用)") } else if len(CORS.Origins) == 0 { log.Printf("ALLOWED_ORIGINS: 未设置(跨域请求将被拒绝)") } else { log.Printf("ALLOWED_ORIGINS: 已配置 %d 个允许的跨域来源白名单", len(CORS.Origins)) } - // 火山 TTS 必填项逐项状态:缺则 ❌,有则 ✓(API Key 脱敏,仅显示头尾各 4 字符) log.Printf("火山 TTS 必填项状态:") type ttsCheck struct { name string @@ -178,15 +259,15 @@ func LogStartupSummary() { ok bool } checks := []ttsCheck{ - {"BYTEDANCE_TTS_API_KEY", maskAPIKey(TTSConfig.ApiKey), TTSConfig.ApiKey != ""}, - {"BYTEDANCE_TTS_RESOURCE_ID", TTSConfig.ResourceId, TTSConfig.ResourceId != ""}, - {"BYTEDANCE_TTS_SPEAKER", TTSConfig.Speaker, TTSConfig.Speaker != ""}, + {"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 = "❌" + mark = "✗" missingCount++ } val := c.value @@ -197,14 +278,12 @@ func LogStartupSummary() { } if TTSConfigErr != nil { - log.Printf("火山 TTS 整体: 初始化失败,%d 个必填项缺失,/v1/audio/speech 路由将全部返回 500", missingCount) + log.Printf("火山 TTS 整体: 初始化失败(%d 个必填项缺失),/v1/audio/speech 路由将全部返回 500", missingCount) } else { log.Printf("火山 TTS 整体: 初始化成功") } } -// maskAPIKey 对 API Key 脱敏,显示头 4 / 尾 4 字符,中间 * 号代替。 -// 短于等于 8 字符整体掩为 ****,空串原样返回。 func maskAPIKey(key string) string { if key == "" { return "" @@ -215,40 +294,13 @@ func maskAPIKey(key string) string { return key[:4] + "****" + key[len(key)-4:] } -// CheckEnvironmentVariables 返回环境变量状态,供 /health 端点使用。 -// 不再直接 os.Getenv,改为读已初始化的全局 Config(单一数据源)。 -func CheckEnvironmentVariables() map[string]interface{} { - requiredVars := map[string]bool{ - "BYTEDANCE_TTS_API_KEY": TTSConfig.ApiKey != "", - "BYTEDANCE_TTS_RESOURCE_ID": TTSConfig.ResourceId != "", - "BYTEDANCE_TTS_SPEAKER": TTSConfig.Speaker != "", - } - - missingVars := []string{} - for varName, isSet := range requiredVars { - if !isSet { - missingVars = append(missingVars, varName) - } - } - - optionalVars := map[string]bool{ - - "OPENAI_TTS_API_KEY": len(Auth.APIKeys) > 0, - "ALLOWED_ORIGINS": CORS.AllowAll || len(CORS.Origins) > 0, - "PORT": Server.Port != common.DefaultPort, - } - - return map[string]interface{}{ - "all_required_vars_set": len(missingVars) == 0, - "missing_required_vars": missingVars, - "required_vars_set": requiredVars, - "optional_vars_set": optionalVars, - } -} - -// CheckStaticFiles 静态文件存在性检查,/dashboard 路由需要 health.html。 +// 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..46be864 --- /dev/null +++ b/telemetry/meter.go @@ -0,0 +1,51 @@ +package telemetry + +import "net/http" + +// Meter 是 telemetry 的高层入口。所有业务埋点都通过它创建指标; +// 测试时可替换为 NoopMeter 或带缓冲的自定义实现。 +type Meter struct { + reg *Registry +} + +// NewMeter 创建默认实现。 +func NewMeter() *Meter { + return &Meter{reg: newRegistry()} +} + +// Handler 返回 /metrics 端点的 http.Handler。 +func (m *Meter) Handler() http.Handler { return m.reg.Handler() } + +// Registry 暴露给特殊用例(如测试断言),生产代码不应使用。 +func (m *Meter) Registry() *Registry { return m.reg } + +// NewCounter 注册并返回一个 Counter。 +// - name 指标名(Prometheus 风格,如 "tts_request_total") +// - help 帮助文本 +// - labelNames 注册时锁定的 label key 集合,运行期不可变 +func (m *Meter) 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 +} + +// NewGauge 同 NewCounter。 +func (m *Meter) 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 +} + +// NewHistogram 同 NewCounter,额外接受桶上界(不含 +Inf,+Inf 由实现自动追加)。 +func (m *Meter) 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..5b06539 --- /dev/null +++ b/telemetry/noop.go @@ -0,0 +1,20 @@ +package telemetry + +import "net/http" + +// NoopMeter 是一个不采集、不输出的 Meter,用于单元测试或禁用观测的场景。 +// 返回的 Counter / Gauge / Histogram 实例不会被注册到任何 Registry, +// 它们的 Inc/Add/Observe 调用在本进程内没有可见效果(每次返回新的空实例)。 +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 +}