feat: WS handler 接入 SessionManager,main.go 注入依赖 + 健康检查获取活跃会话数
Some checks failed
Backend CI / ci (pull_request) Failing after 15s
Frontend CI / ci (pull_request) Has been cancelled

This commit is contained in:
hhs
2026-06-13 15:31:15 +08:00
parent 26f87ce35a
commit fe4dedec71
2 changed files with 78 additions and 36 deletions

View File

@@ -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)
}