179 lines
4.7 KiB
Go
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) })
|
|
}
|