Files
CamTalk/backend/internal/ws/handler.go
hhs fe4dedec71
Some checks failed
Backend CI / ci (pull_request) Failing after 15s
Frontend CI / ci (pull_request) Has been cancelled
feat: WS handler 接入 SessionManager,main.go 注入依赖 + 健康检查获取活跃会话数
2026-06-13 15:31:15 +08:00

183 lines
4.8 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"
"net/http"
"sync"
"time"
"github.com/gin-gonic/gin"
"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{
CheckOrigin: func(r *http.Request) bool { return true }, // 开发阶段允许所有来源
}
// Client 代表一个 WebSocket 客户端连接。
type Client struct {
conn *websocket.Conn
sessionID string
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)
}
// ServeWS 处理 WebSocket 升级请求。
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)
return
}
defer conn.Close()
// 创建会话
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{
Type: "connected",
SessionID: sessionID,
ServerVersion: "0.1.0",
})
logger.Log.Infow("client connected", "session", sessionID)
// 心跳检测
lastPong := time.Now()
conn.SetPongHandler(func(string) error {
lastPong = time.Now()
return nil
})
// 启动心跳检查 goroutine
done := make(chan struct{})
go func() {
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-ticker.C:
if time.Since(lastPong) > 60*time.Second {
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":
_ = 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)
// 刷新会话 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 {
errors.SendWSError(client, errors.CodeInvalidMessage, "", err)
continue
}
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)
// 获取活跃请求 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{
Type: "error",
Code: "INVALID_MESSAGE",
Message: "unknown message type: " + envelope.Type,
})
}
}
close(done)
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
logger.Log.Infow("client disconnected", "session", sessionID)
}