183 lines
4.8 KiB
Go
183 lines
4.8 KiB
Go
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)
|
||
}
|