Files
voice_test/main.go
T

244 lines
6.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"`
}
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
hist []chatMsg
cancel context.CancelFunc
}
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("/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}
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)
case "stop":
s.mu.Lock()
p := s.prov
s.mu.Unlock()
if p != nil {
p.finish()
}
}
case websocket.BinaryMessage:
s.mu.Lock()
p := s.prov
s.mu.Unlock()
if p != nil {
p.sendAudio(data)
}
}
}
}
func (s *session) startASR(id string) {
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 = p, cancel
s.mu.Unlock()
go s.eventPump(p)
log.Printf("ASR started: %s", id)
}
func (s *session) stopASR() {
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()
}
}
// eventPump: ASR 事件 → 转发给浏览器;final 句触发客服回答(LLM 流式 → TTS 流式)。
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.sendJSON(map[string]any{"type": ev.Typ, "text": ev.Text})
if ev.Typ == "final" {
s.agentTurn(p, ev.Text)
}
}
}
}
// agentTurn: final 句 → LLM 流式出字(逐段回传)→ 边生成边按句喂 TTS → 音频流回传。
func (s *session) agentTurn(p asrProvider, userText string) {
ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second)
defer cancel()
s.mu.Lock()
hist := append([]chatMsg(nil), s.hist...)
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, sen, s.sendAudio); 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() {
if sen := strings.TrimSpace(string(pending)); sen != "" {
sentences <- sen
}
pending = pending[:0]
}
for seg := range 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 == ',') {
flush()
}
}
}
flush()
close(sentences)
if err := <-errCh; err != nil {
s.sendJSON(map[string]any{"type": "error", "error": "LLM: " + err.Error()})
<-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"})
}