Files

179 lines
4.7 KiB
Go

package main
import (
"context"
"fmt"
"strings"
"sync"
"sync/atomic"
"github.com/gorilla/websocket"
)
// bailianASR: 阿里云百炼实时语音识别(Qwen-Audio-3.0-ASR-Flash-Streaming / Fun-ASR)。
// 文本事件 run-task/continue-task/finish-task,音频为裸二进制帧。
type bailianASR struct {
model string
api agentCfg
conn *websocket.Conn
taskID string
out chan asrEvent
wmu sync.Mutex // 上游连接写锁(sendAudio 与 updateContext 并发)
connOnce sync.Once
outOnce sync.Once
stopping atomic.Bool
}
func newBailianASR(model string, api agentCfg) *bailianASR {
return &bailianASR{model: model, api: api, out: make(chan asrEvent, 64)}
}
func (a *bailianASR) start(ctx context.Context) error {
conn, _, err := websocket.DefaultDialer.DialContext(ctx, a.api.BailianWssBaseURL,
map[string][]string{"Authorization": {"Bearer " + a.api.BailianKey}})
if err != nil {
return fmt.Errorf("连接百炼 ASR: %w", err)
}
a.conn = conn
started := false
defer func() {
if !started {
conn.Close()
}
}()
a.taskID = uuid()
err = a.writeJSON(map[string]any{
"header": map[string]string{"action": "run-task", "task_id": a.taskID, "streaming": "duplex"},
"payload": map[string]any{
"task_group": "audio", "task": "asr", "function": "recognition",
"model": a.model,
"parameters": map[string]any{"format": "pcm", "sample_rate": 16000},
"input": map[string]any{},
},
})
if err != nil {
return fmt.Errorf("发送 run-task: %w", err)
}
var ev struct {
Header struct {
Event string `json:"event"`
ErrorCode string `json:"error_code"`
ErrorMessage string `json:"error_message"`
} `json:"header"`
}
if err := conn.ReadJSON(&ev); err != nil {
return fmt.Errorf("读取 task-started: %w", err)
}
if ev.Header.Event != "task-started" {
return fmt.Errorf("启动任务失败: %s %s", ev.Header.ErrorCode, ev.Header.ErrorMessage)
}
started = true
go a.readLoop()
return nil
}
func (a *bailianASR) readLoop() {
defer a.closeUpstream()
for {
var ev struct {
Header struct {
Event string `json:"event"`
ErrorCode string `json:"error_code"`
ErrorMessage string `json:"error_message"`
} `json:"header"`
Payload struct {
Output struct {
Sentence struct {
Text string `json:"text"`
SentenceEnd bool `json:"sentence_end"`
} `json:"sentence"`
} `json:"output"`
} `json:"payload"`
}
if err := a.conn.ReadJSON(&ev); err != nil {
if !a.stopping.Load() {
a.emit(asrEvent{Typ: "error", Error: fmt.Sprintf("连接关闭: %v", err)})
}
return
}
switch ev.Header.Event {
case "result-generated":
s := ev.Payload.Output.Sentence
if s.Text == "" {
continue
}
typ := "partial"
if s.SentenceEnd {
typ = "final"
}
a.emit(asrEvent{Typ: typ, Text: s.Text})
case "task-failed":
a.emit(asrEvent{Typ: "error", Code: ev.Header.ErrorCode, Error: ev.Header.ErrorMessage})
return
case "task-finished":
return
}
}
}
// updateContext: 把上一轮对话喂给 ASR(仅支持上下文的模型,fun-asr-flash 系列不支持)。
func (a *bailianASR) updateContext(user, assistant string) error {
if !strings.Contains(a.model, "qwen-audio") && !strings.Contains(a.model, "fun-asr-realtime") {
return nil
}
return a.writeJSON(map[string]any{
"header": map[string]string{"action": "continue-task", "task_id": a.taskID, "streaming": "duplex"},
"payload": map[string]any{"input": map[string]any{"context": []map[string]any{
{"role": "user", "content": []map[string]any{{"type": "input_text", "text": user}}},
{"role": "assistant", "content": []map[string]any{{"type": "text", "text": assistant}}},
}}},
})
}
func (a *bailianASR) emit(e asrEvent) {
select {
case a.out <- e:
default: // ponytail: 前端阻塞时丢帧保命,实时识别不重放
}
}
func (a *bailianASR) writeJSON(v any) error {
a.wmu.Lock()
defer a.wmu.Unlock()
return a.conn.WriteJSON(v)
}
func (a *bailianASR) sendAudio(b []byte) error {
a.wmu.Lock()
defer a.wmu.Unlock()
return a.conn.WriteMessage(websocket.BinaryMessage, b)
}
func (a *bailianASR) finish() error {
return a.writeJSON(map[string]any{
"header": map[string]string{"action": "finish-task", "task_id": a.taskID, "streaming": "duplex"},
"payload": map[string]any{"input": map[string]any{}},
})
}
func (a *bailianASR) events() <-chan asrEvent { return a.out }
func (a *bailianASR) close() {
a.stopping.Store(true)
a.closeConn()
}
func (a *bailianASR) closeConn() {
a.connOnce.Do(func() {
if a.conn != nil {
a.conn.Close()
}
})
}
// 只有 readLoop 退出后才能关闭事件通道,避免 close() 与 emit() 并发导致 panic。
func (a *bailianASR) closeUpstream() {
a.closeConn()
a.outOnce.Do(func() { close(a.out) })
}