diff --git a/backend/internal/ws/handler.go b/backend/internal/ws/handler.go index e91ed4e..6b3500e 100644 --- a/backend/internal/ws/handler.go +++ b/backend/internal/ws/handler.go @@ -15,12 +15,12 @@ import ( "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" + "github.com/hhs/camtalk/internal/trace" ) // newUpgrader 根据配置创建 WebSocket upgrader。 @@ -134,9 +134,20 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche } } + // 生成连接级 trace ID(整个 WebSocket 生命周期使用) + ctx := c.Request.Context() + traceID := trace.GetTraceID(ctx) + if traceID == "" { + // 如果 REST 中间件未生成(不应发生),fallback 生成 + traceID = trace.GenerateTraceID() + ctx = trace.WithTraceID(ctx, traceID) + c.Request = c.Request.WithContext(ctx) + } + conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { - logger.Log.Errorw("websocket upgrade failed", "error", err) + log := trace.FromContext(ctx) + log.Errorw("websocket upgrade failed", "error", err) return } defer conn.Close() @@ -145,13 +156,17 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche var sessionID string if conversationID != "" { sessionID = conversationID - logger.Log.Infow("resuming conversation", "session", sessionID, "user_id", userID) + ctx = trace.WithSessionID(ctx, sessionID) + log := trace.FromContext(ctx) + log.Infow("resuming conversation", "user_id", userID) } else { sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig()) if err != nil { - logger.Log.Errorw("create session failed", "error", err) + log := trace.FromContext(ctx) + log.Errorw("create session failed", "error", err) return } + ctx = trace.WithSessionID(ctx, sessionID) } client := &Client{ @@ -168,7 +183,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche SessionID: sessionID, ServerVersion: version, }) - logger.Log.Infow("client connected", "session", sessionID, "user_id", userID, "username", username) + log := trace.FromContext(ctx) + log.Infow("client connected", "user_id", userID, "username", username) // 心跳检测 lastPong := time.Now() @@ -186,7 +202,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche select { case <-ticker.C: if time.Since(lastPong) > heartbeatTimeout { - logger.Log.Warnw("heartbeat timeout", "session", sessionID) + log := trace.FromContext(ctx) + log.Warnw("heartbeat timeout") conn.Close() return } @@ -201,7 +218,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche _, message, err := conn.ReadMessage() if err != nil { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) { - logger.Log.Warnw("ws read error", "error", err) + log := trace.FromContext(ctx) + log.Warnw("ws read error", "error", err) } break } @@ -226,14 +244,18 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err) continue } - logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID) + + // 注入 request ID 到 context + queryCtx := trace.WithRequestID(ctx, msg.RequestID) + log := trace.FromContext(queryCtx) + log.Infow("query received", "has_image", msg.Image != "", "has_audio", msg.Audio != "") // 限流检查 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) + 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 @@ -242,16 +264,16 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche // 刷新会话 TTL if err := client.sessionMgr.Touch(context.Background(), sessionID); err != nil { - logger.Log.Warnw("touch session failed", "session", sessionID, "error", err) + log.Warnw("touch session failed", "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) + log.Warnw("set active request failed", "error", err) } // 创建可取消的 context - ctx, cancel := context.WithCancel(context.Background()) + processCtx, cancel := context.WithCancel(queryCtx) client.mu.Lock() client.cancelFuncs[msg.RequestID] = cancel client.mu.Unlock() @@ -271,8 +293,9 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche _ = 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) + if err := client.orchestrator.ProcessQuery(processCtx, sessionID, msg, sender); err != nil { + log := trace.FromContext(processCtx) + log.Errorw("process query failed", "error", err) } }() @@ -298,7 +321,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche if msg.Payload.Scenario != nil { scenarioID = *msg.Payload.Scenario } - logger.Log.Infow("config updated", "session", sessionID, "scenario", scenarioID) + log := trace.FromContext(ctx) + log.Infow("config updated", "scenario", scenarioID) // 如果切换了情景(非自由对话),返回首句引导 if scenarioID != "" && scenarioID != "free_chat" { @@ -350,7 +374,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche } case "interrupt": - logger.Log.Infow("interrupt received", "session", sessionID) + log := trace.FromContext(ctx) + log.Infow("interrupt received") // 获取活跃请求 ID 并取消 reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID) @@ -378,12 +403,14 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche // 取消所有活跃请求 client.mu.Lock() for reqID, cancel := range client.cancelFuncs { - logger.Log.Infow("canceling active request on disconnect", "session", sessionID, "request", reqID) + log := trace.FromContext(ctx) + log.Infow("canceling active request on disconnect", "request", reqID) cancel() } client.cancelFuncs = make(map[string]context.CancelFunc) client.mu.Unlock() // 断开连接时不销毁会话,让其自然过期(支持重连恢复) - logger.Log.Infow("client disconnected", "session", sessionID) + log = trace.FromContext(ctx) + log.Infow("client disconnected") }