feat: v2 版本 #206
@@ -15,12 +15,12 @@ import (
|
||||
"github.com/hhs/camtalk/internal/auth"
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/hhs/camtalk/internal/errors"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/orchestrator"
|
||||
"github.com/hhs/camtalk/internal/ratelimit"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// newUpgrader 根据配置创建 WebSocket upgrader。
|
||||
@@ -134,9 +134,20 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
}
|
||||
}
|
||||
|
||||
// 生成连接级 trace ID(整个 WebSocket 生命周期使用)
|
||||
ctx := c.Request.Context()
|
||||
traceID := trace.GetTraceID(ctx)
|
||||
if traceID == "" {
|
||||
// 如果 REST 中间件未生成(不应发生),fallback 生成
|
||||
traceID = trace.GenerateTraceID()
|
||||
ctx = trace.WithTraceID(ctx, traceID)
|
||||
c.Request = c.Request.WithContext(ctx)
|
||||
}
|
||||
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
logger.Log.Errorw("websocket upgrade failed", "error", err)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Errorw("websocket upgrade failed", "error", err)
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
@@ -145,13 +156,17 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
var sessionID string
|
||||
if conversationID != "" {
|
||||
sessionID = conversationID
|
||||
logger.Log.Infow("resuming conversation", "session", sessionID, "user_id", userID)
|
||||
ctx = trace.WithSessionID(ctx, sessionID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Infow("resuming conversation", "user_id", userID)
|
||||
} else {
|
||||
sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig())
|
||||
if err != nil {
|
||||
logger.Log.Errorw("create session failed", "error", err)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Errorw("create session failed", "error", err)
|
||||
return
|
||||
}
|
||||
ctx = trace.WithSessionID(ctx, sessionID)
|
||||
}
|
||||
|
||||
client := &Client{
|
||||
@@ -168,7 +183,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
SessionID: sessionID,
|
||||
ServerVersion: version,
|
||||
})
|
||||
logger.Log.Infow("client connected", "session", sessionID, "user_id", userID, "username", username)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Infow("client connected", "user_id", userID, "username", username)
|
||||
|
||||
// 心跳检测
|
||||
lastPong := time.Now()
|
||||
@@ -186,7 +202,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
select {
|
||||
case <-ticker.C:
|
||||
if time.Since(lastPong) > heartbeatTimeout {
|
||||
logger.Log.Warnw("heartbeat timeout", "session", sessionID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("heartbeat timeout")
|
||||
conn.Close()
|
||||
return
|
||||
}
|
||||
@@ -201,7 +218,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
_, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
|
||||
logger.Log.Warnw("ws read error", "error", err)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("ws read error", "error", err)
|
||||
}
|
||||
break
|
||||
}
|
||||
@@ -226,14 +244,18 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err)
|
||||
continue
|
||||
}
|
||||
logger.Log.Infow("query received", "session", sessionID, "request", msg.RequestID)
|
||||
|
||||
// 注入 request ID 到 context
|
||||
queryCtx := trace.WithRequestID(ctx, msg.RequestID)
|
||||
log := trace.FromContext(queryCtx)
|
||||
log.Infow("query received", "has_image", msg.Image != "", "has_audio", msg.Audio != "")
|
||||
|
||||
// 限流检查
|
||||
if limiter != nil {
|
||||
key := fmt.Sprintf("%s:query", userID)
|
||||
allowed, retryAfter := limiter.Allow(context.Background(), key)
|
||||
if !allowed {
|
||||
logger.Log.Warnw("rate limited", "user_id", userID, "retry_after", retryAfter)
|
||||
log.Warnw("rate limited", "user_id", userID, "retry_after", retryAfter)
|
||||
errors.SendWSError(client, errors.CodeRateLimited, msg.RequestID,
|
||||
fmt.Errorf("rate limited, retry after %s", retryAfter.Round(time.Second)))
|
||||
continue
|
||||
@@ -242,16 +264,16 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
|
||||
// 刷新会话 TTL
|
||||
if err := client.sessionMgr.Touch(context.Background(), sessionID); err != nil {
|
||||
logger.Log.Warnw("touch session failed", "session", sessionID, "error", err)
|
||||
log.Warnw("touch session failed", "error", err)
|
||||
}
|
||||
|
||||
// 标记活跃请求
|
||||
if err := client.sessionMgr.SetActiveRequest(context.Background(), sessionID, msg.RequestID); err != nil {
|
||||
logger.Log.Warnw("set active request failed", "session", sessionID, "error", err)
|
||||
log.Warnw("set active request failed", "error", err)
|
||||
}
|
||||
|
||||
// 创建可取消的 context
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
processCtx, cancel := context.WithCancel(queryCtx)
|
||||
client.mu.Lock()
|
||||
client.cancelFuncs[msg.RequestID] = cancel
|
||||
client.mu.Unlock()
|
||||
@@ -271,8 +293,9 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
|
||||
}()
|
||||
|
||||
if err := client.orchestrator.ProcessQuery(ctx, sessionID, msg, sender); err != nil {
|
||||
logger.Log.Errorw("process query failed", "session", sessionID, "request", msg.RequestID, "error", err)
|
||||
if err := client.orchestrator.ProcessQuery(processCtx, sessionID, msg, sender); err != nil {
|
||||
log := trace.FromContext(processCtx)
|
||||
log.Errorw("process query failed", "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -298,7 +321,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
if msg.Payload.Scenario != nil {
|
||||
scenarioID = *msg.Payload.Scenario
|
||||
}
|
||||
logger.Log.Infow("config updated", "session", sessionID, "scenario", scenarioID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Infow("config updated", "scenario", scenarioID)
|
||||
|
||||
// 如果切换了情景(非自由对话),返回首句引导
|
||||
if scenarioID != "" && scenarioID != "free_chat" {
|
||||
@@ -350,7 +374,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
}
|
||||
|
||||
case "interrupt":
|
||||
logger.Log.Infow("interrupt received", "session", sessionID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Infow("interrupt received")
|
||||
|
||||
// 获取活跃请求 ID 并取消
|
||||
reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID)
|
||||
@@ -378,12 +403,14 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
// 取消所有活跃请求
|
||||
client.mu.Lock()
|
||||
for reqID, cancel := range client.cancelFuncs {
|
||||
logger.Log.Infow("canceling active request on disconnect", "session", sessionID, "request", reqID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Infow("canceling active request on disconnect", "request", reqID)
|
||||
cancel()
|
||||
}
|
||||
client.cancelFuncs = make(map[string]context.CancelFunc)
|
||||
client.mu.Unlock()
|
||||
|
||||
// 断开连接时不销毁会话,让其自然过期(支持重连恢复)
|
||||
logger.Log.Infow("client disconnected", "session", sessionID)
|
||||
log = trace.FromContext(ctx)
|
||||
log.Infow("client disconnected")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user