196 lines
6.4 KiB
Go
196 lines
6.4 KiB
Go
// Package service 提供 transcribe 模块的业务服务
|
||
package service
|
||
|
||
import (
|
||
"bytes"
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
"io"
|
||
"mime/multipart"
|
||
"net/http"
|
||
"path"
|
||
"strings"
|
||
"time"
|
||
|
||
"github.com/echochat/backend/app/transcribe/dao"
|
||
)
|
||
|
||
// OpenAICompatibleClient 实现 OpenAI 兼容协议的 STT 调用
|
||
//
|
||
// 标准端点:POST {base_url}/audio/transcriptions
|
||
// 形式:multipart/form-data
|
||
// - file: 音频文件(mp3/wav/m4a/mp4/...)
|
||
// - model: 模型 code(如 whisper-1 / qwen-audio-asr-flash)
|
||
// - response_format: verbose_json(拿 segments)/ json(仅 text)
|
||
// - language: 可选,提示语言以加速识别
|
||
//
|
||
// 阿里云 DashScope / DeepSeek / OpenAI / 智谱(部分)/ 通义千问 都遵循该协议。
|
||
//
|
||
// 调用超时:默认 5 分钟(覆盖大部分 ≤30min 会议)。超时由调用方通过 ctx 控制更精细。
|
||
type OpenAICompatibleClient struct {
|
||
httpClient *http.Client
|
||
}
|
||
|
||
// NewOpenAICompatibleClient 创建实例
|
||
func NewOpenAICompatibleClient() *OpenAICompatibleClient {
|
||
return &OpenAICompatibleClient{
|
||
httpClient: &http.Client{Timeout: 5 * time.Minute},
|
||
}
|
||
}
|
||
|
||
// TranscribeResult openai 兼容响应(verbose_json 模式)
|
||
//
|
||
// 字段非全列:仅取本服务关心的部分。多余字段被丢弃。
|
||
type TranscribeResult struct {
|
||
Text string `json:"text"`
|
||
Language string `json:"language"`
|
||
Duration float64 `json:"duration"`
|
||
Segments []TranscribeSegment `json:"segments"`
|
||
}
|
||
|
||
// TranscribeSegment 单段时间轴文字(verbose_json)
|
||
type TranscribeSegment struct {
|
||
ID int `json:"id"`
|
||
Start float64 `json:"start"`
|
||
End float64 `json:"end"`
|
||
Text string `json:"text"`
|
||
}
|
||
|
||
// Transcribe 调用 STT
|
||
//
|
||
// 流程:
|
||
// 1. 从 fileURL 下载音频字节(HTTP GET,复用 httpClient timeout)
|
||
// 2. 构造 multipart 请求体
|
||
// 3. POST {base_url}/audio/transcriptions,附 Authorization: Bearer {api_key}
|
||
// 4. 解析 JSON 响应
|
||
//
|
||
// 参数:
|
||
// - fileURL:录制文件公网/内网可达 URL(meeting_recordings.file_url)
|
||
// - language:可选,传 "" 让模型自检
|
||
func (c *OpenAICompatibleClient) Transcribe(ctx context.Context, cfg *dao.STTConfig, fileURL, language string) (*TranscribeResult, error) {
|
||
if cfg == nil {
|
||
return nil, fmt.Errorf("STT config is nil")
|
||
}
|
||
if fileURL == "" {
|
||
return nil, fmt.Errorf("file_url is empty")
|
||
}
|
||
|
||
// 1) 下载音频
|
||
audioBytes, filename, err := c.downloadAudio(ctx, fileURL)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("下载录制文件失败: %w", err)
|
||
}
|
||
|
||
// 2) 拼装 multipart
|
||
body := &bytes.Buffer{}
|
||
writer := multipart.NewWriter(body)
|
||
|
||
filePart, err := writer.CreateFormFile("file", filename)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("创建 multipart file 字段失败: %w", err)
|
||
}
|
||
if _, err := filePart.Write(audioBytes); err != nil {
|
||
return nil, fmt.Errorf("写入 multipart file 内容失败: %w", err)
|
||
}
|
||
if err := writer.WriteField("model", cfg.Model.ModelCode); err != nil {
|
||
return nil, fmt.Errorf("写入 model 字段失败: %w", err)
|
||
}
|
||
if err := writer.WriteField("response_format", "verbose_json"); err != nil {
|
||
return nil, fmt.Errorf("写入 response_format 字段失败: %w", err)
|
||
}
|
||
if language != "" {
|
||
_ = writer.WriteField("language", language)
|
||
}
|
||
if err := writer.Close(); err != nil {
|
||
return nil, fmt.Errorf("关闭 multipart writer 失败: %w", err)
|
||
}
|
||
|
||
// 3) 构造请求
|
||
endpoint := buildTranscriptionEndpoint(cfg.Provider.BaseURL)
|
||
if endpoint == "" {
|
||
return nil, fmt.Errorf("STT provider base_url is empty: provider=%s model=%s", cfg.Provider.ProviderCode, cfg.Model.ModelCode)
|
||
}
|
||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, body)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("构造 STT 请求失败: endpoint=%s provider=%s model=%s: %w", endpoint, cfg.Provider.ProviderCode, cfg.Model.ModelCode, err)
|
||
}
|
||
req.Header.Set("Authorization", "Bearer "+cfg.Key.APIKey)
|
||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||
|
||
resp, err := c.httpClient.Do(req)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("调用 STT 接口失败: endpoint=%s provider=%s model=%s: %w", endpoint, cfg.Provider.ProviderCode, cfg.Model.ModelCode, err)
|
||
}
|
||
defer resp.Body.Close()
|
||
|
||
respBytes, err := io.ReadAll(resp.Body)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("读取 STT 响应失败: %w", err)
|
||
}
|
||
|
||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||
// 截断响应体,避免日志爆炸
|
||
preview := string(respBytes)
|
||
if len(preview) > 300 {
|
||
preview = preview[:300] + "..."
|
||
}
|
||
return nil, fmt.Errorf("STT 接口返回 %d: endpoint=%s provider=%s model=%s body=%s", resp.StatusCode, endpoint, cfg.Provider.ProviderCode, cfg.Model.ModelCode, preview)
|
||
}
|
||
|
||
// 4) 解析
|
||
var result TranscribeResult
|
||
if err := json.Unmarshal(respBytes, &result); err != nil {
|
||
preview := string(respBytes)
|
||
if len(preview) > 300 {
|
||
preview = preview[:300] + "..."
|
||
}
|
||
return nil, fmt.Errorf("解析 STT 响应失败: %v, body=%s", err, preview)
|
||
}
|
||
if result.Text == "" && len(result.Segments) == 0 {
|
||
// 部分供应商在 response_format=json 模式只返回 text;这里没拿到任何结果视为异常
|
||
return nil, fmt.Errorf("STT 响应内容为空")
|
||
}
|
||
return &result, nil
|
||
}
|
||
|
||
func buildTranscriptionEndpoint(baseURL string) string {
|
||
base := strings.TrimRight(strings.TrimSpace(baseURL), "/")
|
||
if base == "" {
|
||
return ""
|
||
}
|
||
if strings.HasSuffix(base, "/audio/transcriptions") {
|
||
return base
|
||
}
|
||
return base + "/audio/transcriptions"
|
||
}
|
||
|
||
// downloadAudio 从给定 URL 拉取音频内容,返回字节流 + 推断的文件名
|
||
//
|
||
// 文件名仅作为 multipart 的 filename 参数,主要决定 Content-Type 推断;
|
||
// 取 URL path 末段,无后缀则默认 recording.webm(与 media-server 实际产物一致)。
|
||
func (c *OpenAICompatibleClient) downloadAudio(ctx context.Context, fileURL string) ([]byte, string, error) {
|
||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, fileURL, nil)
|
||
if err != nil {
|
||
return nil, "", err
|
||
}
|
||
resp, err := c.httpClient.Do(req)
|
||
if err != nil {
|
||
return nil, "", err
|
||
}
|
||
defer resp.Body.Close()
|
||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||
return nil, "", fmt.Errorf("下载失败 status=%d", resp.StatusCode)
|
||
}
|
||
data, err := io.ReadAll(resp.Body)
|
||
if err != nil {
|
||
return nil, "", err
|
||
}
|
||
|
||
filename := path.Base(req.URL.Path)
|
||
if filename == "" || filename == "/" || !strings.Contains(filename, ".") {
|
||
filename = "recording.mp4"
|
||
}
|
||
return data, filename, nil
|
||
}
|