Files
voice_test/main.go
T

318 lines
8.1 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package main
import (
"context"
"encoding/json"
"log"
"net/http"
"strings"
"sync"
"time"
"github.com/gorilla/websocket"
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(*http.Request) bool { return true }, // 本地测试页,放开来源
ReadBufferSize: 16384,
WriteBufferSize: 16384,
}
// asrModels: 前端下拉框的可选模型清单。
func asrModels(cfg agentCfg) []map[string]any {
return []map[string]any{
{"id": "qwen-audio-3.0-asr-flash-streaming", "label": "阿里 Qwen-Audio-3.0-ASR-Flash-Streaming", "vendor": "bailian"},
{"id": "fun-asr-realtime", "label": "阿里 Fun-ASR-Realtime", "vendor": "bailian"},
{"id": "fun-asr-flash-2026-06-15", "label": "阿里 Fun-ASR-Flash (2026-06-15)", "vendor": "bailian"},
{"id": "volc-bigmodel", "label": "火山 Doubao-Seed-ASR-Streaming (sauc)", "vendor": "volc",
"ready": cfg.VolcAppKey != "" || (cfg.VolcAppID != "" && cfg.VolcAccessToken != "")},
}
}
type clientMsg struct {
Type string `json:"type"`
ASR string `json:"asr"`
Mode string `json:"mode"`
TTS *ttsCfg `json:"tts"`
}
type chatMsg struct {
Role string `json:"role"`
Content string `json:"content"`
}
type session struct {
cfg agentCfg
conn *websocket.Conn
wmu sync.Mutex // 客户端连接写锁
mu sync.Mutex // provider / history 状态锁
prov asrProvider
tts ttsCfg
hist []chatMsg
cancel context.CancelFunc
turnCancel context.CancelFunc
turnID uint64
}
func (s *session) sendJSON(v any) {
s.wmu.Lock()
defer s.wmu.Unlock()
s.conn.WriteJSON(v)
}
func (s *session) sendAudio(b []byte) {
s.wmu.Lock()
defer s.wmu.Unlock()
s.conn.WriteMessage(websocket.BinaryMessage, b)
}
func main() {
cfg := loadCfg()
if cfg.BailianKey == "" {
log.Fatal("缺少 BAILIAN_API_KEY 等环境变量(先 source ~/.zshenv")
}
log.Printf("LLM=%s TTS=%s voice=%s", cfg.LLMModel, cfg.TTSModel, cfg.TTSVoice)
http.Handle("/", http.FileServer(http.Dir("web")))
http.HandleFunc("/api/models", func(w http.ResponseWriter, r *http.Request) {
json.NewEncoder(w).Encode(asrModels(cfg))
})
http.HandleFunc("/api/config", func(w http.ResponseWriter, r *http.Request) {
ctx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
defer cancel()
cloned, voiceErr := listClonedVoices(ctx, cfg)
found := false
for _, v := range cloned {
found = found || v.ID == cfg.TTSVoice
}
if cfg.TTSVoice != "" && !found {
cloned = append([]voiceOption{{ID: cfg.TTSVoice, Label: cfg.TTSVoice + "(当前配置)"}}, cloned...)
}
payload := map[string]any{"tts": defaultTTS(cfg), "ttsModels": ttsModels(), "systemVoices": systemVoices(), "clonedVoices": cloned}
if voiceErr != nil {
payload["voiceError"] = voiceErr.Error()
}
json.NewEncoder(w).Encode(payload)
})
http.HandleFunc("/agent", func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
return
}
s := &session{cfg: cfg, conn: conn, tts: defaultTTS(cfg)}
defer conn.Close()
s.readLoop()
})
addr := getenv("PORT", ":8090")
log.Println("listening http://localhost" + addr)
log.Fatal(http.ListenAndServe(addr, nil))
}
// readLoop: 浏览器 → 服务端消息(JSON 控制 / 二进制音频)。
func (s *session) readLoop() {
for {
mt, data, err := s.conn.ReadMessage()
if err != nil {
return
}
switch mt {
case websocket.TextMessage:
var m clientMsg
if json.Unmarshal(data, &m) != nil {
continue
}
switch m.Type {
case "start":
s.startASR(m.ASR, m.TTS)
case "stop":
s.stopASR()
}
case websocket.BinaryMessage:
s.mu.Lock()
p := s.prov
s.mu.Unlock()
if p != nil {
p.sendAudio(data)
}
}
}
}
func (s *session) startASR(id string, requestedTTS *ttsCfg) {
tts, err := normalizeTTS(requestedTTS, s.cfg)
if err != nil {
s.sendJSON(map[string]any{"type": "error", "error": err.Error()})
return
}
s.stopASR()
ctx, cancel := context.WithCancel(context.Background())
var p asrProvider
switch {
case id == "volc-bigmodel":
p = newVolcASR(s.cfg)
default: // 百炼系模型 id 直接透传
p = newBailianASR(id, s.cfg)
}
if err := p.start(ctx); err != nil {
cancel()
s.sendJSON(map[string]any{"type": "error", "error": err.Error()})
return
}
s.mu.Lock()
s.prov, s.cancel, s.tts = p, cancel, tts
s.mu.Unlock()
go s.eventPump(p)
log.Printf("ASR started: %s", id)
}
func (s *session) stopASR() {
s.interruptTurn()
s.mu.Lock()
p, cancel := s.prov, s.cancel
s.prov, s.cancel = nil, nil
s.mu.Unlock()
if cancel != nil {
cancel()
}
if p != nil {
p.finish()
p.close()
}
}
func (s *session) interruptTurn() {
s.mu.Lock()
cancel := s.turnCancel
s.turnCancel = nil
s.mu.Unlock()
if cancel != nil {
cancel()
if s.conn != nil {
s.sendJSON(map[string]string{"type": "interrupt"})
}
}
}
func (s *session) startTurn() (context.Context, context.CancelFunc, uint64) {
ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second)
s.mu.Lock()
s.turnID++
id := s.turnID
s.turnCancel = cancel
s.mu.Unlock()
return ctx, cancel, id
}
func (s *session) finishTurn(id uint64) {
s.mu.Lock()
if s.turnID == id {
s.turnCancel = nil
}
s.mu.Unlock()
}
// eventPump: 新语音一出现即打断当前回答;final 句异步触发下一轮,保持 ASR 事件可继续消费。
func (s *session) eventPump(p asrProvider) {
defer s.sendJSON(map[string]string{"type": "asr-stopped"})
for ev := range p.events() {
switch ev.Typ {
case "error":
s.sendJSON(map[string]any{"type": "error", "error": ev.Error, "code": ev.Code})
case "partial", "final":
s.interruptTurn()
s.sendJSON(map[string]any{"type": ev.Typ, "text": ev.Text})
if ev.Typ == "final" {
ctx, cancel, id := s.startTurn()
go s.agentTurn(ctx, cancel, id, p, ev.Text)
}
}
}
}
// agentTurn: final 句 → LLM 流式出字(逐段回传)→ 边生成边按句喂 TTS → 音频流回传。
func (s *session) agentTurn(ctx context.Context, cancel context.CancelFunc, id uint64, p asrProvider, userText string) {
defer cancel()
defer s.finishTurn(id)
s.mu.Lock()
hist := append([]chatMsg(nil), s.hist...)
tts := s.tts
s.mu.Unlock()
hist = append(hist, chatMsg{Role: "user", Content: userText})
// TTS worker:顺序消费句子,保证音频顺序
sentences := make(chan string, 8)
ttsDone := make(chan struct{})
go func() {
defer close(ttsDone)
for sen := range sentences {
if err := ttsSpeak(ctx, s.cfg, tts, sen, s.sendAudio); err != nil && ctx.Err() == nil {
s.sendJSON(map[string]any{"type": "error", "error": err.Error()})
}
}
}()
textCh, errCh := llmChat(ctx, s.cfg, hist)
var full []byte
var pending []rune
flush := func() bool {
if sen := strings.TrimSpace(string(pending)); sen != "" {
select {
case sentences <- sen:
case <-ctx.Done():
return false
}
}
pending = pending[:0]
return true
}
for seg := range textCh {
if ctx.Err() != nil {
continue // drain producer so cancellation cannot leave it blocked on textCh
}
full = append(full, seg...)
s.sendJSON(map[string]any{"type": "reply-delta", "text": seg})
for _, r := range seg {
pending = append(pending, r)
// 句末标点即切句;超长无标点时退到逗号,避免一句卡住整段
if strings.ContainsRune("。!?!?;\n", r) || (len(pending) >= 100 && r == '') {
if !flush() {
break
}
}
}
}
if ctx.Err() == nil {
flush()
}
close(sentences)
if err := <-errCh; err != nil {
if ctx.Err() == nil {
s.sendJSON(map[string]any{"type": "error", "error": "LLM: " + err.Error()})
}
<-ttsDone
return
}
if ctx.Err() != nil {
<-ttsDone
return
}
reply := string(full)
s.sendJSON(map[string]any{"type": "reply-done", "text": reply})
s.mu.Lock()
s.hist = append(hist, chatMsg{Role: "assistant", Content: reply})
if len(s.hist) > 10 { // ponytail: 内存里只留最近5轮,够客服上下文
s.hist = s.hist[len(s.hist)-10:]
}
s.mu.Unlock()
// qwen streaming 模型支持 continue-task 上下文,提升追问识别
if ba, ok := p.(*bailianASR); ok {
ba.updateContext(userText, reply)
}
<-ttsDone
s.sendJSON(map[string]string{"type": "tts-end"})
}