You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
 
 

375 lines
12 KiB

package filetrans
import (
"encoding/binary"
"encoding/json"
"fmt"
"io"
"net/http"
"sort"
"strings"
"time"
)
// 两套端点/模型:
// - 工作空间模式(配了 WorkspaceId):qwen-audio-3.0-asr-flash-filetrans,端点带 ws 前缀,27 语种;
// - 兼容模式(没配 WorkspaceId):老的 paraformer-v2 + 公共端点,仅 8 语种。
//
// 请求体/响应体两者完全同构(异步 task_id + 轮询 + diarization_enabled + language_hints),
// 所以只有 URL 和 model 名不同,其余解析逻辑共用。
const (
legacyBaseURL = "https://dashscope.aliyuncs.com"
legacyModel = "paraformer-v2"
workspaceBaseURLFmt = "https://%s.cn-beijing.maas.aliyuncs.com"
workspaceModel = "qwen-audio-3.0-asr-flash-filetrans"
pathTranscription = "/api/v1/services/audio/asr/transcription"
pathTasks = "/api/v1/tasks/"
)
// modelLangs 各模型 language_hints 实际支持的语种(短码,即 toDashScopeLang 的输出)。
// 用于 CreateTask 前置校验——把不支持的语种送进去,阿里不报错而是**静默返回空转写**
// (线上出过「西语送进 paraformer-v2 拿到空结果」的事故),所以这里必须 fail-fast,
// 不能只依赖后台 languages 白名单配对。
var modelLangs = map[string]map[string]bool{
legacyModel: {
"zh": true, "en": true, "ja": true, "yue": true,
"ko": true, "de": true, "fr": true, "ru": true,
},
// qwen-audio-3.0-asr-flash-filetrans 官方语种表(中文含普通话/粤语/吴/闽南/客家等方言)。
// 注意:土耳其语 tr、乌克兰语 uk 只出现在 Qwen-ASR 系列的语种表里,本模型的表中没有,
// 未经实测前不放开——宁可漏识别,也不能静默返回空。
workspaceModel: {
"zh": true, "yue": true, "en": true, "ja": true, "ko": true,
"vi": true, "th": true, "id": true, "ms": true, "fil": true,
"hi": true, "ar": true, "fr": true, "de": true, "es": true,
"pt": true, "ru": true, "it": true, "nl": true, "sv": true,
"da": true, "fi": true, "no": true, "el": true, "pl": true,
"cs": true, "hu": true, "ro": true, "bg": true, "hr": true,
"sk": true,
},
}
// model 返回本实例应使用的模型名
func (this *FileTrans) model() string {
if this.options.WorkspaceId != "" {
return workspaceModel
}
return legacyModel
}
// baseURL 返回本实例应使用的服务端点
func (this *FileTrans) baseURL() string {
if this.options.WorkspaceId != "" {
return fmt.Sprintf(workspaceBaseURLFmt, this.options.WorkspaceId)
}
return legacyBaseURL
}
func newSys(options Options) (sys *FileTrans, err error) {
sys = &FileTrans{
options: options,
client: &http.Client{
Timeout: 30 * time.Second,
},
}
return
}
type FileTrans struct {
options Options
client *http.Client
}
// ==================== 提交任务 ====================
// submitRequest DashScope 提交转写任务请求体
type submitRequest struct {
Model string `json:"model"`
Input submitRequestInput `json:"input"`
Parameters submitRequestParams `json:"parameters"`
}
type submitRequestInput struct {
FileURLs []string `json:"file_urls"`
}
type submitRequestParams struct {
LanguageHints []string `json:"language_hints,omitempty"`
SpeakerCount *int `json:"speaker_count,omitempty"`
DiarizationEnabled bool `json:"diarization_enabled"`
NotifyURL string `json:"notify_url,omitempty"`
// ChannelID 要识别的音轨。不传时 DashScope 只转写第 0 轨。
//
// ⚠️ 通话录音是客户端录的**双声道 WAV**(左=本端麦克风、右=对端),
// 不传这个参数的后果就是「语音纪要只转写出我自己说的话,对方一句都没有」。
// 每多一轨按一份音频计费,所以只在文件确实是多声道时才传(见 probeWavChannels)。
ChannelID []int `json:"channel_id,omitempty"`
}
// submitResponse DashScope 提交任务的响应
type submitResponse struct {
RequestID string `json:"request_id"`
Output struct {
TaskID string `json:"task_id"`
TaskStatus string `json:"task_status"`
} `json:"output"`
Code string `json:"code"`
Message string `json:"message"`
}
// ==================== 查询任务 ====================
// queryResponse DashScope 查询任务的响应
type queryResponse struct {
RequestID string `json:"request_id"`
Output struct {
TaskID string `json:"task_id"`
TaskStatus string `json:"task_status"`
Code string `json:"code"`
Message string `json:"message"`
TaskMetrics struct {
Total int `json:"TOTAL"`
Succeeded int `json:"SUCCEEDED"`
Failed int `json:"FAILED"`
} `json:"task_metrics"`
Results []taskResult `json:"results"`
} `json:"output"`
Code string `json:"code"`
Message string `json:"message"`
}
type taskResult struct {
FileURL string `json:"file_url"`
TranscriptionURL string `json:"transcription_url"`
SubtaskStatus string `json:"subtask_status"`
FailureReason string `json:"failure_reason,omitempty"`
}
// transcriptionResult 转写结果 JSON(从 transcription_url 获取)
type transcriptionResult struct {
FileURL string `json:"file_url"`
Properties struct {
AudioFormat string `json:"audio_format"`
Channels []int `json:"channels"`
OriginalSamplingRate int `json:"original_sampling_rate"`
OriginalDurationInMs int64 `json:"original_duration_in_milliseconds"`
} `json:"properties"`
Transcripts []transcript `json:"transcripts"`
}
type transcript struct {
ChannelID int `json:"channel_id"`
Sentences []sentence `json:"sentences"`
}
type sentence struct {
Text string `json:"text"`
BeginTime int64 `json:"begin_time"`
EndTime int64 `json:"end_time"`
SpeakerID any `json:"speaker_id,omitempty"` // 开启说话人分离时为数字,关闭时为字符串
}
// ==================== 实现 ====================
func (this *FileTrans) CreateTask(audioURL, language string, enableSpeaker bool, callbackURL string) (taskID string, err error) {
model := this.model()
params := submitRequestParams{
DiarizationEnabled: enableSpeaker,
NotifyURL: callbackURL,
}
if ch := this.probeWavChannels(audioURL); ch >= 2 {
params.ChannelID = []int{0, 1}
}
if language != "" {
// 不支持的语种绝不放行:阿里对此不报错,只会返回空转写,错误会一路藏到用户看到空白纪要
if !modelLangs[model][language] {
return "", fmt.Errorf("模型 %s 不支持语种 %q,请检查后台该服务的「支持识别的语言」配置", model, language)
}
params.LanguageHints = []string{language}
}
reqBody := submitRequest{
Model: model,
Input: submitRequestInput{
FileURLs: []string{audioURL},
},
Parameters: params,
}
bodyBytes, err := json.Marshal(reqBody)
if err != nil {
return "", fmt.Errorf("marshal request: %w", err)
}
req, err := http.NewRequest("POST", this.baseURL()+pathTranscription, strings.NewReader(string(bodyBytes)))
if err != nil {
return "", fmt.Errorf("new request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+this.options.ApiKey)
req.Header.Set("X-DashScope-Async", "enable")
resp, err := this.client.Do(req)
if err != nil {
return "", fmt.Errorf("do request: %w", err)
}
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
if err != nil {
return "", fmt.Errorf("read response: %w", err)
}
var result submitResponse
if err = json.Unmarshal(data, &result); err != nil {
return "", fmt.Errorf("unmarshal response: %w", err)
}
if result.Code != "" {
return "", fmt.Errorf("create task failed: code=%s, message=%s", result.Code, result.Message)
}
return result.Output.TaskID, nil
}
func (this *FileTrans) QueryTask(taskID string) (status string, contexts []ContextStruct, err error) {
reqURL := this.baseURL() + pathTasks + taskID
req, err := http.NewRequest("GET", reqURL, nil)
if err != nil {
return StatusFailed, nil, fmt.Errorf("new request: %w", err)
}
req.Header.Set("Authorization", "Bearer "+this.options.ApiKey)
resp, err := this.client.Do(req)
if err != nil {
return StatusFailed, nil, fmt.Errorf("do request: %w", err)
}
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
if err != nil {
return StatusFailed, nil, fmt.Errorf("read response: %w", err)
}
var result queryResponse
if err = json.Unmarshal(data, &result); err != nil {
return StatusFailed, nil, fmt.Errorf("unmarshal response: %w", err)
}
if result.Code != "" {
return StatusFailed, nil, fmt.Errorf("query failed: code=%s, message=%s", result.Code, result.Message)
}
switch result.Output.TaskStatus {
case "SUCCEEDED":
contexts, err = this.fetchTranscriptionResults(result.Output.Results)
if err != nil {
return StatusFailed, nil, err
}
return StatusSuccess, contexts, nil
case "FAILED":
reason := result.Output.Message
if reason == "" && len(result.Output.Results) > 0 {
reason = result.Output.Results[0].FailureReason
}
if reason == "" {
reason = result.Message
}
return StatusFailed, nil, fmt.Errorf("task failed: %s", reason)
case "PENDING", "RUNNING":
return StatusRunning, nil, nil
default:
// UNKNOWN / 任务已过期(DashScope 任务结果仅保留 24 小时)/ 其他未知状态
return StatusFailed, nil, fmt.Errorf("task unavailable: task_status=%s output.code=%s output.message=%s top.code=%s top.message=%s raw=%s",
result.Output.TaskStatus, result.Output.Code, result.Output.Message, result.Code, result.Message, string(data))
}
}
// fetchTranscriptionResults 从 transcription_url 下载转写结果并解析
func (this *FileTrans) fetchTranscriptionResults(results []taskResult) (contexts []ContextStruct, err error) {
contexts = make([]ContextStruct, 0)
for _, r := range results {
if r.SubtaskStatus != "SUCCEEDED" || r.TranscriptionURL == "" {
continue
}
resp, err := this.client.Get(r.TranscriptionURL)
if err != nil {
return nil, fmt.Errorf("fetch transcription: %w", err)
}
data, err := io.ReadAll(resp.Body)
resp.Body.Close()
if err != nil {
return nil, fmt.Errorf("read transcription: %w", err)
}
var tr transcriptionResult
if err = json.Unmarshal(data, &tr); err != nil {
return nil, fmt.Errorf("unmarshal transcription: %w", err)
}
// 多音轨(通话录音左右声道)时每轨一份 transcript:说话人直接按音轨给
// (0=本端 1=对端,比分离出来的 speaker_id 可靠得多),并按时间轴合并。
multi := len(tr.Transcripts) > 1
for _, t := range tr.Transcripts {
for _, s := range t.Sentences {
speaker := fmt.Sprintf("%v", s.SpeakerID)
if multi {
speaker = fmt.Sprintf("%d", t.ChannelID)
}
contexts = append(contexts, ContextStruct{
Content: s.Text,
StartTime: s.BeginTime,
EndTime: s.EndTime,
Speaker: speaker,
})
}
}
}
sort.SliceStable(contexts, func(i, j int) bool { return contexts[i].StartTime < contexts[j].StartTime })
return
}
// probeWavChannels 只取文件开头几十个字节,看是不是多声道 WAV。
// 不是 WAV / 取不到 / 解析不出来一律返回 0(按单轨提交,行为与以前一致)。
func (this *FileTrans) probeWavChannels(audioURL string) int {
req, err := http.NewRequest(http.MethodGet, audioURL, nil)
if err != nil {
return 0
}
req.Header.Set("Range", "bytes=0-63")
client := &http.Client{Timeout: 5 * time.Second}
resp, err := client.Do(req)
if err != nil {
return 0
}
defer resp.Body.Close()
head, err := io.ReadAll(io.LimitReader(resp.Body, 64))
if err != nil {
return 0
}
return wavChannelsFromHeader(head)
}
// wavChannelsFromHeader 从 RIFF/WAVE 头里读 fmt 块的声道数;不是 WAV 或不完整返回 0。
func wavChannelsFromHeader(head []byte) int {
if len(head) < 12 || string(head[0:4]) != "RIFF" || string(head[8:12]) != "WAVE" {
return 0
}
// 从 12 起是若干 chunk:4 字节 id + 4 字节长度 + 数据。fmt 块里偏移 2 处是声道数。
pos := 12
for pos+8 <= len(head) {
id := string(head[pos : pos+4])
size := int(binary.LittleEndian.Uint32(head[pos+4 : pos+8]))
if id == "fmt " {
if pos+8+4 > len(head) {
return 0
}
return int(binary.LittleEndian.Uint16(head[pos+10 : pos+12]))
}
pos += 8 + size
}
return 0
}