feat: v2 版本 #206

Merged
huanghaosheng merged 201 commits from v2 into main 2026-06-22 16:01:28 +08:00
Showing only changes of commit 6ab4776e08 - Show all commits

View File

@@ -15,12 +15,12 @@ import (
"github.com/hhs/camtalk/internal/auth" "github.com/hhs/camtalk/internal/auth"
"github.com/hhs/camtalk/internal/config" "github.com/hhs/camtalk/internal/config"
"github.com/hhs/camtalk/internal/errors" "github.com/hhs/camtalk/internal/errors"
"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/orchestrator"
"github.com/hhs/camtalk/internal/ratelimit" "github.com/hhs/camtalk/internal/ratelimit"
"github.com/hhs/camtalk/internal/session" "github.com/hhs/camtalk/internal/session"
"github.com/hhs/camtalk/internal/store" "github.com/hhs/camtalk/internal/store"
"github.com/hhs/camtalk/internal/trace"
) )
// newUpgrader 根据配置创建 WebSocket upgrader。 // 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) conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil { if err != nil {
logger.Log.Errorw("websocket upgrade failed", "error", err) log := trace.FromContext(ctx)
log.Errorw("websocket upgrade failed", "error", err)
return return
} }
defer conn.Close() defer conn.Close()
@@ -145,13 +156,17 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
var sessionID string var sessionID string
if conversationID != "" { if conversationID != "" {
sessionID = 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 { } else {
sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig()) sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig())
if err != nil { if err != nil {
logger.Log.Errorw("create session failed", "error", err) log := trace.FromContext(ctx)
log.Errorw("create session failed", "error", err)
return return
} }
ctx = trace.WithSessionID(ctx, sessionID)
} }
client := &Client{ client := &Client{
@@ -168,7 +183,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
SessionID: sessionID, SessionID: sessionID,
ServerVersion: version, 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() lastPong := time.Now()
@@ -186,7 +202,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
select { select {
case <-ticker.C: case <-ticker.C:
if time.Since(lastPong) > heartbeatTimeout { if time.Since(lastPong) > heartbeatTimeout {
logger.Log.Warnw("heartbeat timeout", "session", sessionID) log := trace.FromContext(ctx)
log.Warnw("heartbeat timeout")
conn.Close() conn.Close()
return return
} }
@@ -201,7 +218,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
_, message, err := conn.ReadMessage() _, message, err := conn.ReadMessage()
if err != nil { if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) { 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 break
} }
@@ -226,14 +244,18 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err) errors.SendWSError(client, errors.CodeInvalidMessage, msg.RequestID, err)
continue 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 { if limiter != nil {
key := fmt.Sprintf("%s:query", userID) key := fmt.Sprintf("%s:query", userID)
allowed, retryAfter := limiter.Allow(context.Background(), key) allowed, retryAfter := limiter.Allow(context.Background(), key)
if !allowed { 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, errors.SendWSError(client, errors.CodeRateLimited, msg.RequestID,
fmt.Errorf("rate limited, retry after %s", retryAfter.Round(time.Second))) fmt.Errorf("rate limited, retry after %s", retryAfter.Round(time.Second)))
continue continue
@@ -242,16 +264,16 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
// 刷新会话 TTL // 刷新会话 TTL
if err := client.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) log.Warnw("touch session failed", "error", err)
} }
// 标记活跃请求 // 标记活跃请求
if err := client.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) log.Warnw("set active request failed", "error", err)
} }
// 创建可取消的 context // 创建可取消的 context
ctx, cancel := context.WithCancel(context.Background()) processCtx, cancel := context.WithCancel(queryCtx)
client.mu.Lock() client.mu.Lock()
client.cancelFuncs[msg.RequestID] = cancel client.cancelFuncs[msg.RequestID] = cancel
client.mu.Unlock() client.mu.Unlock()
@@ -271,8 +293,9 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID) _ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
}() }()
if err := client.orchestrator.ProcessQuery(ctx, sessionID, msg, sender); err != nil { if err := client.orchestrator.ProcessQuery(processCtx, sessionID, msg, sender); err != nil {
logger.Log.Errorw("process query failed", "session", sessionID, "request", msg.RequestID, "error", err) 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 { if msg.Payload.Scenario != nil {
scenarioID = *msg.Payload.Scenario 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" { if scenarioID != "" && scenarioID != "free_chat" {
@@ -350,7 +374,8 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
} }
case "interrupt": case "interrupt":
logger.Log.Infow("interrupt received", "session", sessionID) log := trace.FromContext(ctx)
log.Infow("interrupt received")
// 获取活跃请求 ID 并取消 // 获取活跃请求 ID 并取消
reqID, _ := client.sessionMgr.GetActiveRequestID(context.Background(), sessionID) 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() client.mu.Lock()
for reqID, cancel := range client.cancelFuncs { 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() cancel()
} }
client.cancelFuncs = make(map[string]context.CancelFunc) client.cancelFuncs = make(map[string]context.CancelFunc)
client.mu.Unlock() client.mu.Unlock()
// 断开连接时不销毁会话,让其自然过期(支持重连恢复) // 断开连接时不销毁会话,让其自然过期(支持重连恢复)
logger.Log.Infow("client disconnected", "session", sessionID) log = trace.FromContext(ctx)
log.Infow("client disconnected")
} }