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" ) // 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) 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) } } 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) { // --- 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 { greeting := llm.GetScenarioGreeting(scenarioID, sess.Config.Language) 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) }