Files
CamTalk/backend/internal/ws/handler.go

390 lines
11 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 ws
import (
"context"
"encoding/json"
"fmt"
"net/http"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"github.com/hhs/camtalk/internal/ai/llm"
"github.com/hhs/camtalk/internal/auth"
"github.com/hhs/camtalk/internal/config"
"github.com/hhs/camtalk/internal/errors"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/orchestrator"
"github.com/hhs/camtalk/internal/ratelimit"
"github.com/hhs/camtalk/internal/session"
"github.com/hhs/camtalk/internal/store"
)
// newUpgrader 根据配置创建 WebSocket upgrader。
func newUpgrader(cfg *config.Config) websocket.Upgrader {
allowedOrigins := cfg.Server.AllowedOrigins
return websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
if len(allowedOrigins) == 0 {
return true // 未配置则允许所有来源(开发模式)
}
origin := r.Header.Get("Origin")
for _, o := range allowedOrigins {
if o == origin || o == "*" {
return true
}
}
return false
},
}
}
// Client 代表一个 WebSocket 客户端连接。
type Client struct {
conn *websocket.Conn
sessionID string
sessionMgr session.Manager
orchestrator orchestrator.Orchestrator
cancelFuncs map[string]context.CancelFunc // requestID → cancel func
mu sync.Mutex
}
// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。
func (c *Client) SendJSON(v any) error {
c.mu.Lock()
defer c.mu.Unlock()
return c.conn.WriteJSON(v)
}
// WSClient 实现 orchestrator.Sender 接口,将消息推送到 WebSocket 连接。
type WSClient struct {
client *Client
requestID string
}
// SendSTTResult 发送语音识别结果。
func (w *WSClient) SendSTTResult(result models.WsSTTResult) error {
result.RequestID = w.requestID
return w.client.SendJSON(result)
}
// SendLLMChunk 发送 LLM 流式文本增量。
func (w *WSClient) SendLLMChunk(chunk models.WsLLMChunk) error {
chunk.RequestID = w.requestID
return w.client.SendJSON(chunk)
}
// SendLLMDone 发送 LLM 流结束信号。
func (w *WSClient) SendLLMDone(done models.WsLLMDone) error {
done.RequestID = w.requestID
return w.client.SendJSON(done)
}
// SendTTSAudio 发送 TTS 音频数据。
func (w *WSClient) SendTTSAudio(audio models.WsTTSAudio) error {
audio.RequestID = w.requestID
return w.client.SendJSON(audio)
}
// SendError 发送错误消息。
func (w *WSClient) SendError(err models.WsError) error {
err.RequestID = w.requestID
return w.client.SendJSON(err)
}
// ServeWS 处理 WebSocket 升级请求。
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *config.Config, tokenMgr *auth.TokenManager, limiter ratelimit.Limiter, scenarioRepo store.UserScenarioRepository) gin.HandlerFunc {
upgrader := newUpgrader(cfg)
heartbeatInterval := time.Duration(cfg.Server.HeartbeatInterval) * time.Second
heartbeatTimeout := time.Duration(cfg.Server.HeartbeatTimeout) * time.Second
version := cfg.App.Version
return func(c *gin.Context) {
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, tokenMgr, limiter, scenarioRepo)
}
}
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator,
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, tokenMgr *auth.TokenManager, limiter ratelimit.Limiter, scenarioRepo store.UserScenarioRepository) {
// --- JWT 认证upgrade 前完成,失败直接返回 HTTP 错误) ---
token := c.Query("token")
if token == "" {
c.JSON(http.StatusUnauthorized, gin.H{"error": "missing token"})
return
}
claims, err := tokenMgr.ValidateAccess(token)
if err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid token"})
return
}
userID := claims.UserID
username := claims.Username
// --- conversation_id 处理upgrade 前校验归属) ---
conversationID := c.Query("conversation_id")
if conversationID != "" {
sess, err := sessionMgr.Get(c.Request.Context(), conversationID)
if err != nil || sess.UserID != userID {
c.JSON(http.StatusUnauthorized, gin.H{"error": "SESSION_NOT_FOUND"})
return
}
}
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
logger.Log.Errorw("websocket upgrade failed", "error", err)
return
}
defer conn.Close()
// 创建或复用会话
var sessionID string
if conversationID != "" {
sessionID = conversationID
logger.Log.Infow("resuming conversation", "session", sessionID, "user_id", userID)
} else {
sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig())
if err != nil {
logger.Log.Errorw("create session failed", "error", err)
return
}
}
client := &Client{
conn: conn,
sessionID: sessionID,
sessionMgr: sessionMgr,
orchestrator: orch,
cancelFuncs: make(map[string]context.CancelFunc),
}
// 发送 connected 消息
_ = client.SendJSON(models.WsConnected{
Type: "connected",
SessionID: sessionID,
ServerVersion: version,
})
logger.Log.Infow("client connected", "session", sessionID, "user_id", userID, "username", username)
// 心跳检测
lastPong := time.Now()
conn.SetPongHandler(func(string) error {
lastPong = time.Now()
return nil
})
// 启动心跳检查 goroutine
done := make(chan struct{})
go func() {
ticker := time.NewTicker(heartbeatInterval)
defer ticker.Stop()
for {
select {
case <-ticker.C:
if time.Since(lastPong) > heartbeatTimeout {
logger.Log.Warnw("heartbeat timeout", "session", sessionID)
conn.Close()
return
}
case <-done:
return
}
}
}()
// 消息读取循环
for {
_, message, err := conn.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
logger.Log.Warnw("ws read error", "error", err)
}
break
}
// 解析消息类型
var envelope struct {
Type string `json:"type"`
}
if err := json.Unmarshal(message, &envelope); err != nil {
errors.SendWSError(client, errors.CodeInvalidMessage, "", err)
continue
}
switch envelope.Type {
case "ping":
lastPong = time.Now() // 刷新心跳计时器
_ = client.SendJSON(models.WsPong{Type: "pong"})
case "query":
var msg models.WsQuery
if err := json.Unmarshal(message, &msg); err != nil {
errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err)
continue
}
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
// 限流检查
if limiter != nil {
key := fmt.Sprintf("%s:query", userID)
allowed, retryAfter := limiter.Allow(context.Background(), key)
if !allowed {
logger.Log.Warnw("rate limited", "user_id", userID, "retry_after", retryAfter)
errors.SendWSError(client, errors.CodeRateLimited, msg.RequestID,
fmt.Errorf("rate limited, retry after %s", retryAfter.Round(time.Second)))
continue
}
}
// 刷新会话 TTL
if err := client.sessionMgr.Touch(context.Background(), sessionID); err != nil {
logger.Log.Warnw("touch session failed", "session", sessionID, "error", err)
}
// 标记活跃请求
if err := client.sessionMgr.SetActiveRequest(context.Background(), sessionID, msg.RequestID); err != nil {
logger.Log.Warnw("set active request failed", "session", sessionID, "error", err)
}
// 创建可取消的 context
ctx, cancel := context.WithCancel(context.Background())
client.mu.Lock()
client.cancelFuncs[msg.RequestID] = cancel
client.mu.Unlock()
// 创建 sender
sender := &WSClient{client: client, requestID: msg.RequestID}
// 启动 orchestrator 处理 goroutine
go func() {
defer func() {
// 清理 cancel func
client.mu.Lock()
delete(client.cancelFuncs, msg.RequestID)
client.mu.Unlock()
cancel()
// 清除活跃请求
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
}()
if err := client.orchestrator.ProcessQuery(ctx, sessionID, msg, sender); err != nil {
logger.Log.Errorw("process query failed", "session", sessionID, "request", msg.RequestID, "error", err)
}
}()
case "config":
var msg models.WsConfig
if err := json.Unmarshal(message, &msg); err != nil {
errors.SendWSError(client, errors.CodeInvalidMessage, "", err)
continue
}
patch := models.SessionConfigPatch{
TTSEnabled: msg.Payload.TTSEnabled,
DetailLevel: msg.Payload.DetailLevel,
Language: msg.Payload.Language,
Scenario: msg.Payload.Scenario,
}
if err := client.sessionMgr.UpdateConfig(context.Background(), sessionID, patch); err != nil {
errors.SendWSError(client, errors.CodeInternalError, "", err)
continue
}
scenarioID := ""
if msg.Payload.Scenario != nil {
scenarioID = *msg.Payload.Scenario
}
logger.Log.Infow("config updated", "session", sessionID, "scenario", scenarioID)
// 如果切换了情景(非自由对话),返回首句引导
if scenarioID != "" && scenarioID != "free_chat" {
sess, err := client.sessionMgr.Get(context.Background(), sessionID)
if err == nil && sess != nil {
// 加载用户自建情景
var customGreetings map[string]string
if sess.UserID != "" && scenarioRepo != nil {
scenarios, err := scenarioRepo.FindByUserID(context.Background(), sess.UserID)
if err == nil && len(scenarios) > 0 {
customGreetings = make(map[string]string, len(scenarios))
for _, s := range scenarios {
if s.Greeting != "" {
customGreetings[s.ID] = s.Greeting
}
}
}
}
greeting := llm.GetScenarioGreeting(scenarioID, sess.Config.Language, customGreetings)
if greeting != "" {
// 发送首句作为 AI 消息
_ = client.SendJSON(models.WsLLMChunk{
Type: "llm_chunk",
RequestID: "scenario_greeting",
Delta: greeting,
Role: "assistant",
})
doneMsg := models.WsLLMDone{
Type: "llm_done",
RequestID: "scenario_greeting",
FullText: greeting,
Model: "",
LatencyMs: 0,
}
doneMsg.TokensUsed.Prompt = 0
doneMsg.TokensUsed.Completion = 0
doneMsg.TokensUsed.Total = 0
_ = client.SendJSON(doneMsg)
// 追加首句到历史记录
_ = client.sessionMgr.AppendMessage(context.Background(), sessionID, models.Message{
Role: "assistant",
Content: greeting,
})
}
}
}
case "interrupt":
logger.Log.Infow("interrupt received", "session", sessionID)
// 获取活跃请求 ID 并取消
reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID)
if reqID != "" {
client.mu.Lock()
if cancel, ok := client.cancelFuncs[reqID]; ok {
cancel()
delete(client.cancelFuncs, reqID)
}
client.mu.Unlock()
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
}
default:
_ = client.SendJSON(models.WsError{
Type: "error",
Code: "INVALID_MESSAGE",
Message: "unknown message type: " + envelope.Type,
})
}
}
close(done)
// 取消所有活跃请求
client.mu.Lock()
for reqID, cancel := range client.cancelFuncs {
logger.Log.Infow("canceling active request on disconnect", "session", sessionID, "request", reqID)
cancel()
}
client.cancelFuncs = make(map[string]context.CancelFunc)
client.mu.Unlock()
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
logger.Log.Infow("client disconnected", "session", sessionID)
}