diff --git a/backend/internal/ws/handler.go b/backend/internal/ws/handler.go index 8b98e38..feff0ac 100644 --- a/backend/internal/ws/handler.go +++ b/backend/internal/ws/handler.go @@ -122,6 +122,16 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche 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) @@ -129,11 +139,17 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche } defer conn.Close() - // 创建会话 - sessionID, err := sessionMgr.Create(context.Background(), userID, models.DefaultConfig()) - if err != nil { - logger.Log.Errorw("create session failed", "error", err) - return + // 创建或复用会话 + 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{