diff --git a/backend/internal/ai/stt/deepgram.go b/backend/internal/ai/stt/deepgram.go new file mode 100644 index 0000000..94ca970 --- /dev/null +++ b/backend/internal/ai/stt/deepgram.go @@ -0,0 +1,139 @@ +package stt + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/url" + "strings" + "time" + + "github.com/gorilla/websocket" + "go.uber.org/zap" +) + +// DeepgramService 基于 Deepgram WebSocket API 的语音识别实现。 +type DeepgramService struct { + apiKey string + endpoint string + logger *zap.SugaredLogger +} + +// NewDeepgramService 创建 Deepgram STT 服务。 +func NewDeepgramService(apiKey, endpoint string, logger *zap.SugaredLogger) *DeepgramService { + if endpoint == "" { + endpoint = "wss://api.deepgram.com/v1/listen" + } + return &DeepgramService{ + apiKey: apiKey, + endpoint: endpoint, + logger: logger, + } +} + +// deepgramResponse Deepgram WebSocket 响应。 +type deepgramResponse struct { + Channel struct { + Alternatives []struct { + Transcript string `json:"transcript"` + Confidence float64 `json:"confidence"` + } `json:"alternatives"` + } `json:"channel"` + IsFinal bool `json:"is_final"` +} + +// Recognize 实现 stt.Service。通过 WebSocket 发送音频到 Deepgram,返回最终识别文本。 +func (d *DeepgramService) Recognize(ctx context.Context, audio []byte, opts Options) (string, error) { + if len(audio) == 0 { + return "", fmt.Errorf("stt: empty audio") + } + + // 构建 WebSocket URL,附带查询参数 + wsURL := d.buildURL(opts) + + // 5 秒总超时 + ctx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + + // 建立 WebSocket 连接 + conn, _, err := websocket.DefaultDialer.DialContext(ctx, wsURL, http.Header{ + "Authorization": []string{"Token " + d.apiKey}, + }) + if err != nil { + return "", fmt.Errorf("stt: connect deepgram: %w", err) + } + defer conn.Close() + + // 发送音频数据(一次性) + if err := conn.WriteMessage(websocket.BinaryMessage, audio); err != nil { + return "", fmt.Errorf("stt: send audio: %w", err) + } + + // 发送 Close 消息通知服务端音频已发送完毕 + closeMsg := websocket.FormatCloseMessage(websocket.CloseNormalClosure, "") + _ = conn.WriteMessage(websocket.CloseMessage, closeMsg) + + // 读取识别结果 + var transcript strings.Builder + for { + _, message, err := conn.ReadMessage() + if err != nil { + // Close 帧是正常的结束信号 + if websocket.IsCloseError(err, websocket.CloseNormalClosure) { + break + } + if websocket.IsUnexpectedCloseError(err, websocket.CloseNormalClosure) { + break + } + return "", fmt.Errorf("stt: read response: %w", err) + } + + var resp deepgramResponse + if err := json.Unmarshal(message, &resp); err != nil { + d.logger.Warnw("stt: unmarshal response failed", "error", err) + continue + } + + if len(resp.Channel.Alternatives) > 0 { + text := strings.TrimSpace(resp.Channel.Alternatives[0].Transcript) + if text != "" { + transcript.WriteString(text) + } + } + + if resp.IsFinal { + break + } + } + + return strings.TrimSpace(transcript.String()), nil +} + +// buildURL 构建 Deepgram WebSocket URL,包含音频格式参数。 +func (d *DeepgramService) buildURL(opts Options) string { + u, _ := url.Parse(d.endpoint) + + encoding := opts.Encoding + if encoding == "" { + encoding = "pcm_s16le" + } + sampleRate := opts.SampleRate + if sampleRate == 0 { + sampleRate = 16000 + } + language := opts.Language + if language == "" { + language = "zh-CN" + } + + q := u.Query() + q.Set("encoding", encoding) + q.Set("sample_rate", fmt.Sprintf("%d", sampleRate)) + q.Set("language", language) + q.Set("model", "nova-2") + q.Set("punctuate", "true") + u.RawQuery = q.Encode() + + return u.String() +}