321 lines
8.2 KiB
Go
321 lines
8.2 KiB
Go
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() {
|
||
if err := loadDotEnv(".env"); err != nil {
|
||
log.Fatalf("加载 .env 失败: %v", err)
|
||
}
|
||
cfg := loadCfg()
|
||
if cfg.BailianKey == "" {
|
||
log.Fatal("缺少 BAILIAN_API_KEY(请配置 .env 或系统环境变量)")
|
||
}
|
||
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"})
|
||
}
|