feat: WS handler 接入 SessionManager,main.go 注入依赖 + 健康检查获取活跃会话数
This commit is contained in:
@@ -12,6 +12,7 @@ import (
|
|||||||
|
|
||||||
"github.com/hhs/camtalk/internal/config"
|
"github.com/hhs/camtalk/internal/config"
|
||||||
"github.com/hhs/camtalk/internal/logger"
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
|
"github.com/hhs/camtalk/internal/session"
|
||||||
"github.com/hhs/camtalk/internal/ws"
|
"github.com/hhs/camtalk/internal/ws"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -33,6 +34,12 @@ func main() {
|
|||||||
"addr", cfg.Server.Addr(),
|
"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 模式
|
// Gin 模式
|
||||||
if cfg.App.Env == "prod" {
|
if cfg.App.Env == "prod" {
|
||||||
gin.SetMode(gin.ReleaseMode)
|
gin.SetMode(gin.ReleaseMode)
|
||||||
@@ -44,11 +51,11 @@ func main() {
|
|||||||
// REST API
|
// REST API
|
||||||
api := r.Group("/api")
|
api := r.Group("/api")
|
||||||
{
|
{
|
||||||
api.GET("/health", healthHandler)
|
api.GET("/health", healthHandler(sessionMgr))
|
||||||
}
|
}
|
||||||
|
|
||||||
// WebSocket
|
// WebSocket
|
||||||
r.GET("/ws", ws.ServeWS)
|
r.GET("/ws", ws.ServeWS(sessionMgr))
|
||||||
|
|
||||||
// HTTP Server
|
// HTTP Server
|
||||||
srv := &http.Server{
|
srv := &http.Server{
|
||||||
@@ -82,11 +89,13 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// healthHandler 健康检查。
|
// healthHandler 健康检查。
|
||||||
func healthHandler(c *gin.Context) {
|
func healthHandler(sessionMgr session.Manager) gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
c.JSON(200, gin.H{
|
c.JSON(200, gin.H{
|
||||||
"status": "ok",
|
"status": "ok",
|
||||||
"version": "0.1.0",
|
"version": "0.1.0",
|
||||||
"uptime": time.Since(startTime).String(),
|
"uptime": time.Since(startTime).String(),
|
||||||
"active_sessions": 0, // TODO: 接入 Session Manager
|
"active_sessions": sessionMgr.ActiveCount(),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,17 +1,19 @@
|
|||||||
package ws
|
package ws
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/google/uuid"
|
|
||||||
"github.com/gorilla/websocket"
|
"github.com/gorilla/websocket"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/errors"
|
||||||
"github.com/hhs/camtalk/internal/logger"
|
"github.com/hhs/camtalk/internal/logger"
|
||||||
"github.com/hhs/camtalk/internal/models"
|
"github.com/hhs/camtalk/internal/models"
|
||||||
|
"github.com/hhs/camtalk/internal/session"
|
||||||
)
|
)
|
||||||
|
|
||||||
var upgrader = websocket.Upgrader{
|
var upgrader = websocket.Upgrader{
|
||||||
@@ -25,14 +27,21 @@ type Client struct {
|
|||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) sendJSON(v any) error {
|
// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。
|
||||||
|
func (c *Client) SendJSON(v any) error {
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
return c.conn.WriteJSON(v)
|
return c.conn.WriteJSON(v)
|
||||||
}
|
}
|
||||||
|
|
||||||
// ServeWS 处理 WebSocket 升级请求。
|
// 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)
|
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Log.Errorw("websocket upgrade failed", "error", err)
|
logger.Log.Errorw("websocket upgrade failed", "error", err)
|
||||||
@@ -40,11 +49,17 @@ func ServeWS(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
defer conn.Close()
|
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}
|
client := &Client{conn: conn, sessionID: sessionID}
|
||||||
|
|
||||||
// 发送 connected 消息
|
// 发送 connected 消息
|
||||||
_ = client.sendJSON(models.WsConnected{
|
_ = client.SendJSON(models.WsConnected{
|
||||||
Type: "connected",
|
Type: "connected",
|
||||||
SessionID: sessionID,
|
SessionID: sessionID,
|
||||||
ServerVersion: "0.1.0",
|
ServerVersion: "0.1.0",
|
||||||
@@ -92,51 +107,67 @@ func ServeWS(c *gin.Context) {
|
|||||||
Type string `json:"type"`
|
Type string `json:"type"`
|
||||||
}
|
}
|
||||||
if err := json.Unmarshal(message, &envelope); err != nil {
|
if err := json.Unmarshal(message, &envelope); err != nil {
|
||||||
_ = client.sendJSON(models.WsError{
|
errors.SendWSError(client, errors.CodeInvalidMessage, "", err)
|
||||||
Type: "error",
|
|
||||||
Code: "INVALID_MESSAGE",
|
|
||||||
Message: "invalid JSON",
|
|
||||||
})
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
switch envelope.Type {
|
switch envelope.Type {
|
||||||
case "ping":
|
case "ping":
|
||||||
_ = client.sendJSON(models.WsPong{Type: "pong"})
|
_ = client.SendJSON(models.WsPong{Type: "pong"})
|
||||||
|
|
||||||
case "query":
|
case "query":
|
||||||
var msg models.WsQuery
|
var msg models.WsQuery
|
||||||
if err := json.Unmarshal(message, &msg); err != nil {
|
if err := json.Unmarshal(message, &msg); err != nil {
|
||||||
_ = client.sendJSON(models.WsError{
|
errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err)
|
||||||
Type: "error",
|
|
||||||
Code: "INVALID_MESSAGE",
|
|
||||||
Message: "invalid query message",
|
|
||||||
RequestID: msg.RequestID,
|
|
||||||
})
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
|
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":
|
case "config":
|
||||||
var msg models.WsConfig
|
var msg models.WsConfig
|
||||||
if err := json.Unmarshal(message, &msg); err != nil {
|
if err := json.Unmarshal(message, &msg); err != nil {
|
||||||
_ = client.sendJSON(models.WsError{
|
errors.SendWSError(client, errors.CodeInvalidMessage, "", err)
|
||||||
Type: "error",
|
|
||||||
Code: "INVALID_MESSAGE",
|
|
||||||
Message: "invalid config message",
|
|
||||||
})
|
|
||||||
continue
|
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":
|
case "interrupt":
|
||||||
logger.Log.Infow("interrupt received", "session", sessionID)
|
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:
|
default:
|
||||||
_ = client.sendJSON(models.WsError{
|
_ = client.SendJSON(models.WsError{
|
||||||
Type: "error",
|
Type: "error",
|
||||||
Code: "INVALID_MESSAGE",
|
Code: "INVALID_MESSAGE",
|
||||||
Message: "unknown message type: " + envelope.Type,
|
Message: "unknown message type: " + envelope.Type,
|
||||||
@@ -145,5 +176,7 @@ func ServeWS(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
close(done)
|
close(done)
|
||||||
|
|
||||||
|
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
|
||||||
logger.Log.Infow("client disconnected", "session", sessionID)
|
logger.Log.Infow("client disconnected", "session", sessionID)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user