feat: 将 Session Manager + Orchestrator 串入 WebSocket handler,实现端到端消息处理。 #38
@@ -10,8 +10,12 @@ import (
|
|||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/hhs/camtalk/internal/ai/llm"
|
||||||
|
"github.com/hhs/camtalk/internal/ai/stt"
|
||||||
|
"github.com/hhs/camtalk/internal/ai/tts"
|
||||||
"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/orchestrator"
|
||||||
"github.com/hhs/camtalk/internal/session"
|
"github.com/hhs/camtalk/internal/session"
|
||||||
"github.com/hhs/camtalk/internal/ws"
|
"github.com/hhs/camtalk/internal/ws"
|
||||||
)
|
)
|
||||||
@@ -40,6 +44,14 @@ func main() {
|
|||||||
sessionMgr = session.NewMemoryManager(30*time.Minute, 20)
|
sessionMgr = session.NewMemoryManager(30*time.Minute, 20)
|
||||||
defer sessionMgr.(*session.MemoryManager).Stop()
|
defer sessionMgr.(*session.MemoryManager).Stop()
|
||||||
|
|
||||||
|
// 初始化 AI 服务
|
||||||
|
sttService := stt.NewDeepgramService(cfg.AI.STT.APIKey, cfg.AI.STT.Endpoint, logger.Log)
|
||||||
|
llmService := llm.NewOpenAIService(cfg.AI.LLM.APIKey, cfg.AI.LLM.Model, cfg.AI.LLM.Endpoint, cfg.AI.LLM.Timeout, logger.Log)
|
||||||
|
ttsService := tts.NewOpenAIService(cfg.AI.TTS.APIKey, cfg.AI.TTS.Voice, cfg.AI.TTS.Endpoint, cfg.AI.TTS.Speed, cfg.AI.TTS.Timeout, logger.Log)
|
||||||
|
|
||||||
|
// 初始化 Orchestrator
|
||||||
|
orch := orchestrator.New(sttService, llmService, ttsService, sessionMgr)
|
||||||
|
|
||||||
// Gin 模式
|
// Gin 模式
|
||||||
if cfg.App.Env == "prod" {
|
if cfg.App.Env == "prod" {
|
||||||
gin.SetMode(gin.ReleaseMode)
|
gin.SetMode(gin.ReleaseMode)
|
||||||
@@ -55,7 +67,7 @@ func main() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// WebSocket
|
// WebSocket
|
||||||
r.GET("/ws", ws.ServeWS(sessionMgr))
|
r.GET("/ws", ws.ServeWS(sessionMgr, orch))
|
||||||
|
|
||||||
// HTTP Server
|
// HTTP Server
|
||||||
srv := &http.Server{
|
srv := &http.Server{
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"github.com/hhs/camtalk/internal/errors"
|
"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/orchestrator"
|
||||||
"github.com/hhs/camtalk/internal/session"
|
"github.com/hhs/camtalk/internal/session"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -22,9 +23,12 @@ var upgrader = websocket.Upgrader{
|
|||||||
|
|
||||||
// Client 代表一个 WebSocket 客户端连接。
|
// Client 代表一个 WebSocket 客户端连接。
|
||||||
type Client struct {
|
type Client struct {
|
||||||
conn *websocket.Conn
|
conn *websocket.Conn
|
||||||
sessionID string
|
sessionID string
|
||||||
mu sync.Mutex
|
sessionMgr session.Manager
|
||||||
|
orchestrator orchestrator.Orchestrator
|
||||||
|
cancelFuncs map[string]context.CancelFunc // requestID → cancel func
|
||||||
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。
|
// SendJSON 向客户端发送 JSON 消息(公开以便 errors 包调用)。
|
||||||
@@ -34,14 +38,50 @@ func (c *Client) SendJSON(v any) error {
|
|||||||
return c.conn.WriteJSON(v)
|
return c.conn.WriteJSON(v)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WSClient 实现 orchestrator.Sender 接口,将消息推送到 WebSocket 连接。
|
||||||
|
type WSClient struct {
|
||||||
|
client *Client
|
||||||
|
requestID string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendSTTResult 发送语音识别结果。
|
||||||
|
func (w *WSClient) SendSTTResult(result models.WsSTTResult) error {
|
||||||
|
result.RequestID = w.requestID
|
||||||
|
return w.client.SendJSON(result)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendLLMChunk 发送 LLM 流式文本增量。
|
||||||
|
func (w *WSClient) SendLLMChunk(chunk models.WsLLMChunk) error {
|
||||||
|
chunk.RequestID = w.requestID
|
||||||
|
return w.client.SendJSON(chunk)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendLLMDone 发送 LLM 流结束信号。
|
||||||
|
func (w *WSClient) SendLLMDone(done models.WsLLMDone) error {
|
||||||
|
done.RequestID = w.requestID
|
||||||
|
return w.client.SendJSON(done)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendTTSAudio 发送 TTS 音频数据。
|
||||||
|
func (w *WSClient) SendTTSAudio(audio models.WsTTSAudio) error {
|
||||||
|
audio.RequestID = w.requestID
|
||||||
|
return w.client.SendJSON(audio)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendError 发送错误消息。
|
||||||
|
func (w *WSClient) SendError(err models.WsError) error {
|
||||||
|
err.RequestID = w.requestID
|
||||||
|
return w.client.SendJSON(err)
|
||||||
|
}
|
||||||
|
|
||||||
// ServeWS 处理 WebSocket 升级请求。
|
// ServeWS 处理 WebSocket 升级请求。
|
||||||
func ServeWS(sessionMgr session.Manager) gin.HandlerFunc {
|
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
serveWS(c, sessionMgr)
|
serveWS(c, sessionMgr, orch)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func serveWS(c *gin.Context, sessionMgr session.Manager) {
|
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator) {
|
||||||
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)
|
||||||
@@ -56,7 +96,13 @@ func serveWS(c *gin.Context, sessionMgr session.Manager) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
client := &Client{conn: conn, sessionID: sessionID}
|
client := &Client{
|
||||||
|
conn: conn,
|
||||||
|
sessionID: sessionID,
|
||||||
|
sessionMgr: sessionMgr,
|
||||||
|
orchestrator: orch,
|
||||||
|
cancelFuncs: make(map[string]context.CancelFunc),
|
||||||
|
}
|
||||||
|
|
||||||
// 发送 connected 消息
|
// 发送 connected 消息
|
||||||
_ = client.SendJSON(models.WsConnected{
|
_ = client.SendJSON(models.WsConnected{
|
||||||
@@ -124,19 +170,43 @@ func serveWS(c *gin.Context, sessionMgr session.Manager) {
|
|||||||
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
|
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
|
||||||
|
|
||||||
// 刷新会话 TTL
|
// 刷新会话 TTL
|
||||||
if err := sessionMgr.Touch(context.Background(), sessionID); err != nil {
|
if err := client.sessionMgr.Touch(context.Background(), sessionID); err != nil {
|
||||||
logger.Log.Warnw("touch session failed", "session", sessionID, "error", err)
|
logger.Log.Warnw("touch session failed", "session", sessionID, "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 标记活跃请求
|
// 标记活跃请求
|
||||||
if err := sessionMgr.SetActiveRequest(context.Background(), sessionID, msg.RequestID); err != nil {
|
if err := client.sessionMgr.SetActiveRequest(context.Background(), sessionID, msg.RequestID); err != nil {
|
||||||
logger.Log.Warnw("set active request failed", "session", sessionID, "error", err)
|
logger.Log.Warnw("set active request failed", "session", sessionID, "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 获取对话历史(供后续 Orchestrator 使用)
|
// 获取对话历史
|
||||||
_, _ = sessionMgr.GetHistory(context.Background(), sessionID, 20)
|
history, _ := client.sessionMgr.GetHistory(context.Background(), sessionID, 20)
|
||||||
|
|
||||||
// TODO: 解码 audio Base64 → 启动 orchestrator.ProcessQuery goroutine
|
// 创建可取消的 context
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
client.mu.Lock()
|
||||||
|
client.cancelFuncs[msg.RequestID] = cancel
|
||||||
|
client.mu.Unlock()
|
||||||
|
|
||||||
|
// 创建 sender
|
||||||
|
sender := &WSClient{client: client, requestID: msg.RequestID}
|
||||||
|
|
||||||
|
// 启动 orchestrator 处理 goroutine
|
||||||
|
go func() {
|
||||||
|
defer func() {
|
||||||
|
// 清理 cancel func
|
||||||
|
client.mu.Lock()
|
||||||
|
delete(client.cancelFuncs, msg.RequestID)
|
||||||
|
client.mu.Unlock()
|
||||||
|
cancel()
|
||||||
|
// 清除活跃请求
|
||||||
|
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
|
||||||
|
}()
|
||||||
|
|
||||||
|
if err := client.orchestrator.ProcessQuery(ctx, sessionID, msg, history, sender); err != nil {
|
||||||
|
logger.Log.Errorw("process query failed", "session", sessionID, "request", msg.RequestID, "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
case "config":
|
case "config":
|
||||||
var msg models.WsConfig
|
var msg models.WsConfig
|
||||||
@@ -150,7 +220,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager) {
|
|||||||
DetailLevel: msg.Payload.DetailLevel,
|
DetailLevel: msg.Payload.DetailLevel,
|
||||||
Language: msg.Payload.Language,
|
Language: msg.Payload.Language,
|
||||||
}
|
}
|
||||||
if err := sessionMgr.UpdateConfig(context.Background(), sessionID, patch); err != nil {
|
if err := client.sessionMgr.UpdateConfig(context.Background(), sessionID, patch); err != nil {
|
||||||
errors.SendWSError(client, errors.CodeInternalError, "", err)
|
errors.SendWSError(client, errors.CodeInternalError, "", err)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -159,11 +229,16 @@ func serveWS(c *gin.Context, sessionMgr session.Manager) {
|
|||||||
case "interrupt":
|
case "interrupt":
|
||||||
logger.Log.Infow("interrupt received", "session", sessionID)
|
logger.Log.Infow("interrupt received", "session", sessionID)
|
||||||
|
|
||||||
// 获取活跃请求 ID(实际 cancel 在 Phase 5 接入 orchestrator 后实现)
|
// 获取活跃请求 ID 并取消
|
||||||
reqID, _ := sessionMgr.GetActiveRequestID(context.Background(), sessionID)
|
reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID)
|
||||||
if reqID != "" {
|
if reqID != "" {
|
||||||
_ = sessionMgr.ClearActiveRequest(context.Background(), sessionID)
|
client.mu.Lock()
|
||||||
// TODO: 取消对应 context cancel func
|
if cancel, ok := client.cancelFuncs[reqID]; ok {
|
||||||
|
cancel()
|
||||||
|
delete(client.cancelFuncs, reqID)
|
||||||
|
}
|
||||||
|
client.mu.Unlock()
|
||||||
|
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
|
||||||
}
|
}
|
||||||
|
|
||||||
default:
|
default:
|
||||||
@@ -177,6 +252,15 @@ func serveWS(c *gin.Context, sessionMgr session.Manager) {
|
|||||||
|
|
||||||
close(done)
|
close(done)
|
||||||
|
|
||||||
|
// 取消所有活跃请求
|
||||||
|
client.mu.Lock()
|
||||||
|
for reqID, cancel := range client.cancelFuncs {
|
||||||
|
logger.Log.Infow("canceling active request on disconnect", "session", sessionID, "request", reqID)
|
||||||
|
cancel()
|
||||||
|
}
|
||||||
|
client.cancelFuncs = make(map[string]context.CancelFunc)
|
||||||
|
client.mu.Unlock()
|
||||||
|
|
||||||
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
|
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
|
||||||
logger.Log.Infow("client disconnected", "session", sessionID)
|
logger.Log.Infow("client disconnected", "session", sessionID)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user