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"}) }