diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index 794ff53..a475fa4 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -12,6 +12,7 @@ import ( "github.com/hhs/camtalk/internal/config" "github.com/hhs/camtalk/internal/logger" + "github.com/hhs/camtalk/internal/session" "github.com/hhs/camtalk/internal/ws" ) @@ -33,6 +34,12 @@ func main() { "addr", cfg.Server.Addr(), ) + // 初始化 Session Manager(MVP 默认内存实现) + var sessionMgr session.Manager + // TODO: 当 Redis 配置非空时切换为 RedisManager + sessionMgr = session.NewMemoryManager(30*time.Minute, 20) + defer sessionMgr.(*session.MemoryManager).Stop() + // Gin 模式 if cfg.App.Env == "prod" { gin.SetMode(gin.ReleaseMode) @@ -44,11 +51,11 @@ func main() { // REST API api := r.Group("/api") { - api.GET("/health", healthHandler) + api.GET("/health", healthHandler(sessionMgr)) } // WebSocket - r.GET("/ws", ws.ServeWS) + r.GET("/ws", ws.ServeWS(sessionMgr)) // HTTP Server srv := &http.Server{ @@ -82,11 +89,13 @@ func main() { } // healthHandler 健康检查。 -func healthHandler(c *gin.Context) { - c.JSON(200, gin.H{ - "status": "ok", - "version": "0.1.0", - "uptime": time.Since(startTime).String(), - "active_sessions": 0, // TODO: 接入 Session Manager - }) +func healthHandler(sessionMgr session.Manager) gin.HandlerFunc { + return func(c *gin.Context) { + c.JSON(200, gin.H{ + "status": "ok", + "version": "0.1.0", + "uptime": time.Since(startTime).String(), + "active_sessions": sessionMgr.ActiveCount(), + }) + } } diff --git a/backend/internal/ws/handler.go b/backend/internal/ws/handler.go index e27f4c0..83316b4 100644 --- a/backend/internal/ws/handler.go +++ b/backend/internal/ws/handler.go @@ -1,17 +1,19 @@ package ws import ( + "context" "encoding/json" "net/http" "sync" "time" "github.com/gin-gonic/gin" - "github.com/google/uuid" "github.com/gorilla/websocket" + "github.com/hhs/camtalk/internal/errors" "github.com/hhs/camtalk/internal/logger" "github.com/hhs/camtalk/internal/models" + "github.com/hhs/camtalk/internal/session" ) var upgrader = websocket.Upgrader{ @@ -25,14 +27,21 @@ type Client struct { mu sync.Mutex } -func (c *Client) sendJSON(v any) error { +// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。 +func (c *Client) SendJSON(v any) error { c.mu.Lock() defer c.mu.Unlock() return c.conn.WriteJSON(v) } // ServeWS 处理 WebSocket 升级请求。 -func ServeWS(c *gin.Context) { +func ServeWS(sessionMgr session.Manager) gin.HandlerFunc { + return func(c *gin.Context) { + serveWS(c, sessionMgr) + } +} + +func serveWS(c *gin.Context, sessionMgr session.Manager) { conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { logger.Log.Errorw("websocket upgrade failed", "error", err) @@ -40,11 +49,17 @@ func ServeWS(c *gin.Context) { } defer conn.Close() - sessionID := uuid.New().String() + // 创建会话 + sessionID, err := sessionMgr.Create(context.Background(), models.DefaultConfig()) + if err != nil { + logger.Log.Errorw("create session failed", "error", err) + return + } + client := &Client{conn: conn, sessionID: sessionID} // 发送 connected 消息 - _ = client.sendJSON(models.WsConnected{ + _ = client.SendJSON(models.WsConnected{ Type: "connected", SessionID: sessionID, ServerVersion: "0.1.0", @@ -92,51 +107,67 @@ func ServeWS(c *gin.Context) { Type string `json:"type"` } if err := json.Unmarshal(message, &envelope); err != nil { - _ = client.sendJSON(models.WsError{ - Type: "error", - Code: "INVALID_MESSAGE", - Message: "invalid JSON", - }) + errors.SendWSError(client, errors.CodeInvalidMessage, "", err) continue } switch envelope.Type { case "ping": - _ = client.sendJSON(models.WsPong{Type: "pong"}) + _ = client.SendJSON(models.WsPong{Type: "pong"}) case "query": var msg models.WsQuery if err := json.Unmarshal(message, &msg); err != nil { - _ = client.sendJSON(models.WsError{ - Type: "error", - Code: "INVALID_MESSAGE", - Message: "invalid query message", - RequestID: msg.RequestID, - }) + errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err) continue } logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID) - // TODO: 调用 AI 编排流程(STT → LLM → TTS) + + // 刷新会话 TTL + if err := sessionMgr.Touch(context.Background(), sessionID); err != nil { + logger.Log.Warnw("touch session failed", "session", sessionID, "error", err) + } + + // 标记活跃请求 + if err := sessionMgr.SetActiveRequest(context.Background(), sessionID, msg.RequestID); err != nil { + logger.Log.Warnw("set active request failed", "session", sessionID, "error", err) + } + + // 获取对话历史(供后续 Orchestrator 使用) + _, _ = sessionMgr.GetHistory(context.Background(), sessionID, 20) + + // TODO: 解码 audio Base64 → 启动 orchestrator.ProcessQuery goroutine case "config": var msg models.WsConfig if err := json.Unmarshal(message, &msg); err != nil { - _ = client.sendJSON(models.WsError{ - Type: "error", - Code: "INVALID_MESSAGE", - Message: "invalid config message", - }) + errors.SendWSError(client, errors.CodeInvalidMessage, "", err) continue } - logger.Log.Infow("config update", "session", sessionID) - // TODO: 更新会话配置 + + patch := models.SessionConfigPatch{ + TTSEnabled: msg.Payload.TTSEnabled, + DetailLevel: msg.Payload.DetailLevel, + Language: msg.Payload.Language, + } + if err := sessionMgr.UpdateConfig(context.Background(), sessionID, patch); err != nil { + errors.SendWSError(client, errors.CodeInternalError, "", err) + continue + } + logger.Log.Infow("config updated", "session", sessionID) case "interrupt": logger.Log.Infow("interrupt received", "session", sessionID) - // TODO: 中断当前 AI 响应 + + // 获取活跃请求 ID(实际 cancel 在 Phase 5 接入 orchestrator 后实现) + reqID, _ := sessionMgr.GetActiveRequestID(context.Background(), sessionID) + if reqID != "" { + _ = sessionMgr.ClearActiveRequest(context.Background(), sessionID) + // TODO: 取消对应 context cancel func + } default: - _ = client.sendJSON(models.WsError{ + _ = client.SendJSON(models.WsError{ Type: "error", Code: "INVALID_MESSAGE", Message: "unknown message type: " + envelope.Type, @@ -145,5 +176,7 @@ func ServeWS(c *gin.Context) { } close(done) + + // 断开连接时不销毁会话,让其自然过期(支持重连恢复) logger.Log.Infow("client disconnected", "session", sessionID) }