feat: 实现日志追踪链路 #185
@@ -23,6 +23,7 @@ import (
|
||||
"github.com/hhs/camtalk/internal/ratelimit"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/hhs/camtalk/internal/ws"
|
||||
migrations "github.com/hhs/camtalk/migrations"
|
||||
)
|
||||
@@ -226,7 +227,9 @@ func main() {
|
||||
}
|
||||
|
||||
r := gin.New()
|
||||
r.Use(gin.Recovery())
|
||||
r.Use(trace.TraceMiddleware()) // 第一层:生成 trace ID
|
||||
r.Use(trace.GinLogger()) // 第二层:记录请求
|
||||
r.Use(trace.GinRecovery()) // 第三层:panic 恢复
|
||||
|
||||
// REST API
|
||||
apiGroup := r.Group("/api")
|
||||
|
||||
@@ -53,6 +53,7 @@ require (
|
||||
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
|
||||
github.com/modern-go/reflect2 v1.0.2 // indirect
|
||||
github.com/nikolalohinski/gonja v1.5.3 // indirect
|
||||
github.com/oklog/ulid/v2 v2.1.1 // indirect
|
||||
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
||||
github.com/pkg/errors v0.9.1 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
|
||||
@@ -129,9 +129,12 @@ github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9G
|
||||
github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk=
|
||||
github.com/nikolalohinski/gonja v1.5.3 h1:GsA+EEaZDZPGJ8JtpeGN78jidhOlxeJROpqMT9fTj9c=
|
||||
github.com/nikolalohinski/gonja v1.5.3/go.mod h1:RmjwxNiXAEqcq1HeK5SSMmqFJvKOfTfXhkJv6YBtPa4=
|
||||
github.com/oklog/ulid/v2 v2.1.1 h1:suPZ4ARWLOJLegGFiZZ1dFAkqzhMjL3J1TzI+5wHz8s=
|
||||
github.com/oklog/ulid/v2 v2.1.1/go.mod h1:rcEKHmBBKfef9DhnvX7y1HZBYxjXb0cP5ExxNsTT1QQ=
|
||||
github.com/onsi/ginkgo v1.6.0/go.mod h1:lLunBs/Ym6LB5Z9jYTR76FiuTmxDTDusOGeTQH+WWjE=
|
||||
github.com/onsi/ginkgo v1.8.0/go.mod h1:lLunBs/Ym6LB5Z9jYTR76FiuTmxDTDusOGeTQH+WWjE=
|
||||
github.com/onsi/gomega v1.5.0/go.mod h1:ex+gbHU/CVuBBDIJjb2X0qEXbFg53c61hWP/1CpauHY=
|
||||
github.com/pborman/getopt v0.0.0-20170112200414-7148bc3a4c30/go.mod h1:85jBQOZwpVEaDAr341tbn15RS4fCAsIst0qp7i8ex1o=
|
||||
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
|
||||
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
|
||||
github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||
|
||||
@@ -11,6 +11,8 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/hhs/camtalk/internal/util"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
@@ -108,7 +110,11 @@ func (m *MiMoService) SynthesizeStream(ctx context.Context, textStream <-chan st
|
||||
|
||||
audio, err := m.synthesize(ctx, text, voice)
|
||||
if err != nil {
|
||||
m.logger.Warnw("mimo tts: synthesize failed", "error", err, "text", text)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("mimo tts: synthesize failed",
|
||||
"error", err,
|
||||
"text_len", len(text),
|
||||
"text_preview", util.Truncate(text, 100))
|
||||
// 静默跳过,不中断整个流
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -9,6 +9,8 @@ import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/hhs/camtalk/internal/util"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
@@ -81,7 +83,11 @@ func (o *OpenAIService) SynthesizeStream(ctx context.Context, textStream <-chan
|
||||
|
||||
audio, err := o.synthesize(ctx, text, voice, speed)
|
||||
if err != nil {
|
||||
o.logger.Warnw("tts: synthesize failed", "error", err, "text", text)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("tts: synthesize failed",
|
||||
"error", err,
|
||||
"text_len", len(text),
|
||||
"text_preview", util.Truncate(text, 100))
|
||||
// 静默跳过,不中断整个流
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"github.com/hhs/camtalk/internal/auth"
|
||||
apperr "github.com/hhs/camtalk/internal/errors"
|
||||
"github.com/hhs/camtalk/internal/ratelimit"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// AuthHandler 提供认证相关的 REST 端点。
|
||||
@@ -53,6 +54,9 @@ func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup, limiter ratelimit.Limi
|
||||
|
||||
// Register POST /api/auth/register — 用户注册。
|
||||
func (h *AuthHandler) Register(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
clientIP := c.ClientIP()
|
||||
|
||||
var req auth.RegisterRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
@@ -72,15 +76,25 @@ func (h *AuthHandler) Register(c *gin.Context) {
|
||||
|
||||
resp, err := h.authService.Register(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
log.Warnw("register failed",
|
||||
"username", req.Username,
|
||||
"client_ip", clientIP,
|
||||
"error", err)
|
||||
handleAuthError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("register success",
|
||||
"username", req.Username,
|
||||
"client_ip", clientIP)
|
||||
c.JSON(http.StatusCreated, resp)
|
||||
}
|
||||
|
||||
// Login POST /api/auth/login — 用户登录。
|
||||
func (h *AuthHandler) Login(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
clientIP := c.ClientIP()
|
||||
|
||||
var req auth.LoginRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
@@ -100,15 +114,24 @@ func (h *AuthHandler) Login(c *gin.Context) {
|
||||
|
||||
resp, err := h.authService.Login(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
log.Warnw("login failed",
|
||||
"username", req.Username,
|
||||
"client_ip", clientIP,
|
||||
"error", err)
|
||||
handleAuthError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("login success",
|
||||
"username", req.Username,
|
||||
"client_ip", clientIP)
|
||||
c.JSON(http.StatusOK, resp)
|
||||
}
|
||||
|
||||
// Refresh POST /api/auth/refresh — 刷新令牌。
|
||||
func (h *AuthHandler) Refresh(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
|
||||
var req auth.RefreshRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
@@ -128,15 +151,21 @@ func (h *AuthHandler) Refresh(c *gin.Context) {
|
||||
|
||||
resp, err := h.authService.Refresh(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
log.Warnw("token refresh failed",
|
||||
"client_ip", c.ClientIP(),
|
||||
"error", err)
|
||||
handleAuthError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("token refresh success",
|
||||
"client_ip", c.ClientIP())
|
||||
c.JSON(http.StatusOK, resp)
|
||||
}
|
||||
|
||||
// Logout POST /api/auth/logout — 登出(需要认证)。
|
||||
func (h *AuthHandler) Logout(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
userID := c.GetString(auth.ContextKeyUserID)
|
||||
|
||||
var req struct {
|
||||
@@ -159,6 +188,9 @@ func (h *AuthHandler) Logout(c *gin.Context) {
|
||||
}
|
||||
|
||||
if err := h.authService.Logout(c.Request.Context(), userID, req.RefreshToken); err != nil {
|
||||
log.Errorw("logout failed",
|
||||
"user_id", userID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to logout",
|
||||
@@ -166,6 +198,8 @@ func (h *AuthHandler) Logout(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("logout success",
|
||||
"user_id", userID)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"message": "logged out successfully",
|
||||
})
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// ConversationHandler 提供对话相关的 REST 端点。
|
||||
@@ -47,6 +48,7 @@ func (h *ConversationHandler) RegisterRoutes(rg *gin.RouterGroup) {
|
||||
|
||||
// List GET /api/conversations — 获取当前用户的对话列表。
|
||||
func (h *ConversationHandler) List(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
userID := c.GetString(auth.ContextKeyUserID)
|
||||
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
@@ -61,6 +63,9 @@ func (h *ConversationHandler) List(c *gin.Context) {
|
||||
|
||||
summaries, total, err := h.sessionMgr.ListByUser(c.Request.Context(), userID, page, size)
|
||||
if err != nil {
|
||||
log.Errorw("list conversations failed",
|
||||
"user_id", userID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to list conversations",
|
||||
@@ -83,6 +88,7 @@ type CreateConversationRequest struct {
|
||||
|
||||
// Create POST /api/conversations — 创建新对话。
|
||||
func (h *ConversationHandler) Create(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
userID := c.GetString(auth.ContextKeyUserID)
|
||||
|
||||
var req CreateConversationRequest
|
||||
@@ -95,6 +101,9 @@ func (h *ConversationHandler) Create(c *gin.Context) {
|
||||
|
||||
sessionID, err := h.sessionMgr.Create(c.Request.Context(), userID, cfg)
|
||||
if err != nil {
|
||||
log.Errorw("create conversation failed",
|
||||
"user_id", userID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to create conversation",
|
||||
@@ -104,6 +113,9 @@ func (h *ConversationHandler) Create(c *gin.Context) {
|
||||
|
||||
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
|
||||
if err != nil {
|
||||
log.Errorw("retrieve created conversation failed",
|
||||
"session_id", sessionID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to retrieve created conversation",
|
||||
@@ -111,6 +123,9 @@ func (h *ConversationHandler) Create(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("conversation created",
|
||||
"conversation_id", sess.ID,
|
||||
"user_id", userID)
|
||||
c.JSON(http.StatusCreated, gin.H{
|
||||
"id": sess.ID,
|
||||
"title": sess.Title,
|
||||
@@ -144,6 +159,7 @@ type UpdateTitleRequest struct {
|
||||
|
||||
// UpdateTitle PATCH /api/conversations/:id — 更新对话标题。
|
||||
func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
sessionID := c.Param("id")
|
||||
|
||||
// 先校验归属
|
||||
@@ -176,6 +192,9 @@ func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
|
||||
})
|
||||
return
|
||||
}
|
||||
log.Errorw("update title failed",
|
||||
"session_id", sessionID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to update title",
|
||||
@@ -190,6 +209,7 @@ func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
|
||||
|
||||
// Delete DELETE /api/conversations/:id — 删除对话。
|
||||
func (h *ConversationHandler) Delete(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
sessionID := c.Param("id")
|
||||
|
||||
// 先校验归属
|
||||
@@ -205,6 +225,9 @@ func (h *ConversationHandler) Delete(c *gin.Context) {
|
||||
})
|
||||
return
|
||||
}
|
||||
log.Errorw("delete conversation failed",
|
||||
"session_id", sessionID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to delete conversation",
|
||||
@@ -221,6 +244,7 @@ func (h *ConversationHandler) Delete(c *gin.Context) {
|
||||
// - limit: 返回消息数量上限,默认 50
|
||||
// - before: 消息 ID 游标(用于分页),返回此 ID 之前的消息
|
||||
func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
sessionID := c.Param("id")
|
||||
|
||||
// 先校验归属
|
||||
@@ -239,6 +263,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
||||
if h.msgRepo != nil {
|
||||
messages, err := h.msgRepo.GetMessages(c.Request.Context(), sessionID, limit, beforeID)
|
||||
if err != nil {
|
||||
log.Errorw("get messages failed",
|
||||
"session_id", sessionID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to get messages",
|
||||
@@ -263,6 +290,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
||||
})
|
||||
return
|
||||
}
|
||||
log.Errorw("get messages failed",
|
||||
"session_id", sessionID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to get messages",
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// SessionHandler 提供会话相关的 REST 端点。
|
||||
@@ -27,6 +28,8 @@ type CreateSessionRequest struct {
|
||||
|
||||
// CreateSession POST /api/sessions — 创建新会话。
|
||||
func (h *SessionHandler) CreateSession(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
|
||||
var req CreateSessionRequest
|
||||
// 请求体可选,解析失败不报错(使用默认配置)
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
@@ -38,6 +41,8 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
||||
|
||||
sessionID, err := h.sessionMgr.Create(c.Request.Context(), "", cfg)
|
||||
if err != nil {
|
||||
log.Errorw("create session failed",
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": "INTERNAL_ERROR",
|
||||
"message": "failed to create session",
|
||||
@@ -48,6 +53,9 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
||||
// 获取创建后的会话以返回 created_at
|
||||
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
|
||||
if err != nil {
|
||||
log.Errorw("retrieve created session failed",
|
||||
"session_id", sessionID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": "INTERNAL_ERROR",
|
||||
"message": "failed to retrieve created session",
|
||||
@@ -55,6 +63,8 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("session created",
|
||||
"session_id", sess.ID)
|
||||
c.JSON(http.StatusCreated, gin.H{
|
||||
"session_id": sess.ID,
|
||||
"created_at": sess.CreatedAt,
|
||||
@@ -63,6 +73,7 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
||||
|
||||
// DestroySession DELETE /api/sessions/:id — 销毁会话。
|
||||
func (h *SessionHandler) DestroySession(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
sessionID := c.Param("id")
|
||||
|
||||
err := h.sessionMgr.Destroy(c.Request.Context(), sessionID)
|
||||
@@ -74,6 +85,9 @@ func (h *SessionHandler) DestroySession(c *gin.Context) {
|
||||
})
|
||||
return
|
||||
}
|
||||
log.Errorw("destroy session failed",
|
||||
"session_id", sessionID,
|
||||
"error", err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": "INTERNAL_ERROR",
|
||||
"message": "failed to destroy session",
|
||||
@@ -81,6 +95,8 @@ func (h *SessionHandler) DestroySession(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("session destroyed",
|
||||
"session_id", sessionID)
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// contextKey 用于在 Gin context 中存储 Claims 的 key。
|
||||
@@ -17,8 +19,13 @@ const (
|
||||
// 校验成功后将 user_id 和 username 写入 Gin Context。
|
||||
func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if authHeader == "" {
|
||||
log.Warnw("auth rejected",
|
||||
"client_ip", c.ClientIP(),
|
||||
"path", c.Request.URL.Path,
|
||||
"reason", "missing authorization header")
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
||||
"code": "INVALID_TOKEN",
|
||||
"message": "missing authorization header",
|
||||
@@ -29,6 +36,10 @@ func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
|
||||
// 提取 Bearer token
|
||||
parts := strings.SplitN(authHeader, " ", 2)
|
||||
if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
|
||||
log.Warnw("auth rejected",
|
||||
"client_ip", c.ClientIP(),
|
||||
"path", c.Request.URL.Path,
|
||||
"reason", "invalid authorization format")
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
||||
"code": "INVALID_TOKEN",
|
||||
"message": "invalid authorization format",
|
||||
@@ -38,6 +49,11 @@ func AuthMiddleware(tokenMgr *TokenManager) gin.HandlerFunc {
|
||||
|
||||
claims, err := tokenMgr.ValidateAccess(parts[1])
|
||||
if err != nil {
|
||||
log.Warnw("auth rejected",
|
||||
"client_ip", c.ClientIP(),
|
||||
"path", c.Request.URL.Path,
|
||||
"reason", "invalid or expired token",
|
||||
"error", err)
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
|
||||
"code": "INVALID_TOKEN",
|
||||
"message": "invalid or expired token",
|
||||
|
||||
@@ -8,20 +8,12 @@ import (
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/orchestrator"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// ctxKeySessionID sessionID 的 context key。
|
||||
type ctxKeySessionID struct{}
|
||||
|
||||
// WithSessionID 将 sessionID 注入 context。
|
||||
func WithSessionID(ctx context.Context, sessionID string) context.Context {
|
||||
return context.WithValue(ctx, ctxKeySessionID{}, sessionID)
|
||||
}
|
||||
|
||||
// EinoOrchestrator 实现 orchestrator.Orchestrator 接口。
|
||||
// 将 Eino Graph 包装为现有接口,WS Handler 几乎不用改。
|
||||
type EinoOrchestrator struct {
|
||||
@@ -48,19 +40,19 @@ func (e *EinoOrchestrator) ProcessQuery(
|
||||
req models.WsQuery,
|
||||
sender orchestrator.Sender,
|
||||
) error {
|
||||
log := logger.Log
|
||||
log := trace.FromContext(ctx)
|
||||
startTime := time.Now()
|
||||
|
||||
// 1. 设置活跃请求
|
||||
if err := e.sessionMgr.SetActiveRequest(ctx, sessionID, req.RequestID); err != nil {
|
||||
log.Errorw("设置活跃请求失败", "error", err)
|
||||
return err
|
||||
}
|
||||
defer e.sessionMgr.ClearActiveRequest(ctx, sessionID)
|
||||
|
||||
// 2. 获取会话配置
|
||||
sess, err := e.sessionMgr.Get(ctx, sessionID)
|
||||
if err != nil {
|
||||
log.Errorw("获取会话失败", "error", err)
|
||||
log.Errorw("get session failed", "error", err)
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
@@ -75,7 +67,7 @@ func (e *EinoOrchestrator) ProcessQuery(
|
||||
if req.Text == "" && req.Audio != "" {
|
||||
audioData, err = base64.StdEncoding.DecodeString(req.Audio)
|
||||
if err != nil {
|
||||
log.Errorw("音频解码失败", "error", err)
|
||||
log.Errorw("audio decode failed", "error", err)
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
@@ -90,7 +82,7 @@ func (e *EinoOrchestrator) ProcessQuery(
|
||||
if req.Image != "" {
|
||||
imageData, err = base64.StdEncoding.DecodeString(req.Image)
|
||||
if err != nil {
|
||||
log.Errorw("图片解码失败", "error", err)
|
||||
log.Errorw("image decode failed", "error", err)
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
@@ -107,7 +99,7 @@ func (e *EinoOrchestrator) ProcessQuery(
|
||||
// 5. 注入 context 值(供 Callback 和 Lambda 节点使用)
|
||||
ctx = WithSender(ctx, sender)
|
||||
ctx = WithRequestID(ctx, req.RequestID)
|
||||
ctx = WithSessionID(ctx, sessionID)
|
||||
ctx = trace.WithSessionID(ctx, sessionID)
|
||||
ctx = WithStartTime(ctx, startTime)
|
||||
|
||||
// 创建 State 并从 input 复制元数据
|
||||
@@ -125,7 +117,7 @@ func (e *EinoOrchestrator) ProcessQuery(
|
||||
// 6. 调用 Graph(Stream 模式 + 运行时 Callback)
|
||||
streamReader, err := e.graph.Runnable.Stream(ctx, input, e.callbacks)
|
||||
if err != nil {
|
||||
log.Errorw("Graph Stream 启动失败", "error", err)
|
||||
log.Errorw("graph stream start failed", "error", err)
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
RequestID: req.RequestID,
|
||||
@@ -143,7 +135,7 @@ func (e *EinoOrchestrator) ProcessQuery(
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
log.Errorw("Graph Stream 消费错误", "error", err)
|
||||
log.Errorw("graph stream consume error", "error", err)
|
||||
break
|
||||
}
|
||||
output = o
|
||||
@@ -159,7 +151,7 @@ func (e *EinoOrchestrator) ProcessQuery(
|
||||
Role: "user",
|
||||
Content: userText,
|
||||
}); err != nil {
|
||||
log.Errorw("追加用户消息到历史失败", "session", sessionID, "error", err)
|
||||
log.Errorw("append user message failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -169,15 +161,12 @@ func (e *EinoOrchestrator) ProcessQuery(
|
||||
Role: "assistant",
|
||||
Content: output.FullResponse,
|
||||
}); err != nil {
|
||||
log.Errorw("追加助手消息到历史失败", "session", sessionID, "error", err)
|
||||
log.Errorw("append assistant message failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
latency := time.Since(startTime).Milliseconds()
|
||||
log.Infow("Eino 编排完成",
|
||||
"request_id", req.RequestID,
|
||||
"latency_ms", latency,
|
||||
"session_id", sessionID)
|
||||
log.Infow("eino pipeline completed", "latency_ms", latency)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -9,14 +9,13 @@ import (
|
||||
"github.com/cloudwego/eino/schema"
|
||||
callbacksHelper "github.com/cloudwego/eino/utils/callbacks"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/orchestrator"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// context key 类型,避免与其他包冲突。
|
||||
type ctxKeySender struct{}
|
||||
type ctxKeyRequestID struct{}
|
||||
type ctxKeyState struct{}
|
||||
|
||||
// WithSender 将 Sender 注入 context。
|
||||
@@ -24,9 +23,9 @@ func WithSender(ctx context.Context, sender orchestrator.Sender) context.Context
|
||||
return context.WithValue(ctx, ctxKeySender{}, sender)
|
||||
}
|
||||
|
||||
// WithRequestID 将 requestID 注入 context。
|
||||
// WithRequestID 将 requestID 注入 context(使用 trace 包)。
|
||||
func WithRequestID(ctx context.Context, requestID string) context.Context {
|
||||
return context.WithValue(ctx, ctxKeyRequestID{}, requestID)
|
||||
return trace.WithRequestID(ctx, requestID)
|
||||
}
|
||||
|
||||
// WithPipelineState 将 PipelineState 注入 context。
|
||||
@@ -40,10 +39,9 @@ func senderFromCtx(ctx context.Context) orchestrator.Sender {
|
||||
return s
|
||||
}
|
||||
|
||||
// requestIDFromCtx 从 context 获取 requestID。
|
||||
// requestIDFromCtx 从 context 获取 requestID(使用 trace 包)。
|
||||
func requestIDFromCtx(ctx context.Context) string {
|
||||
s, _ := ctx.Value(ctxKeyRequestID{}).(string)
|
||||
return s
|
||||
return trace.GetRequestID(ctx)
|
||||
}
|
||||
|
||||
// stateFromCtx 从 context 获取 PipelineState。
|
||||
@@ -62,7 +60,7 @@ func BuildCallbackHandler() callbacks.Handler {
|
||||
return callbacksHelper.NewHandlerHelper().
|
||||
ChatModel(&callbacksHelper.ModelCallbackHandler{
|
||||
OnEndWithStreamOutput: func(ctx context.Context, info *callbacks.RunInfo, output *schema.StreamReader[*model.CallbackOutput]) context.Context {
|
||||
log := logger.Log
|
||||
log := trace.FromContext(ctx)
|
||||
sender := senderFromCtx(ctx)
|
||||
requestID := requestIDFromCtx(ctx)
|
||||
state := stateFromCtx(ctx)
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"github.com/hhs/camtalk/internal/ai/tts"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/orchestrator"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// --- Mock STT Service ---
|
||||
@@ -176,7 +177,7 @@ func TestContextInjection(t *testing.T) {
|
||||
sender := &mockSender{}
|
||||
ctx = WithSender(ctx, sender)
|
||||
ctx = WithRequestID(ctx, "req-123")
|
||||
ctx = WithSessionID(ctx, "sess-456")
|
||||
ctx = trace.WithSessionID(ctx, "sess-456")
|
||||
ctx = WithStartTime(ctx, time.Now())
|
||||
ctx = WithPipelineState(ctx, genLocalState(ctx))
|
||||
|
||||
|
||||
@@ -6,8 +6,8 @@ import (
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// ctxKeyStartTime 请求开始时间的 context key。
|
||||
@@ -33,7 +33,7 @@ func latencyFromCtx(ctx context.Context) int64 {
|
||||
// 历史消息追加由适配器负责(避免重复写入)。
|
||||
func NewDoneLambda(defaultModel string) *compose.Lambda {
|
||||
return compose.InvokableLambda(func(ctx context.Context, _ struct{}) (PipelineOutput, error) {
|
||||
log := logger.Log
|
||||
log := trace.FromContext(ctx)
|
||||
sender := senderFromCtx(ctx)
|
||||
state := stateFromCtx(ctx)
|
||||
|
||||
@@ -70,13 +70,11 @@ func NewDoneLambda(defaultModel string) *compose.Lambda {
|
||||
}
|
||||
}
|
||||
if err := sender.SendLLMDone(done); err != nil {
|
||||
log.Errorw("发送 llm_done 失败", "error", err)
|
||||
log.Errorw("send llm_done failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
log.Infow("查询处理完成",
|
||||
"request_id", requestID,
|
||||
"response_length", len(fullResponse))
|
||||
log.Infow("query processing completed", "response_length", len(fullResponse))
|
||||
|
||||
return PipelineOutput{
|
||||
TranscribedText: transcribedText,
|
||||
|
||||
@@ -8,9 +8,9 @@ import (
|
||||
"github.com/cloudwego/eino/schema"
|
||||
|
||||
"github.com/hhs/camtalk/internal/ai/llm"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// NewHistoryLambda 创建历史组装 Lambda 节点。
|
||||
@@ -24,7 +24,7 @@ func NewHistoryLambda(
|
||||
maxHistory int,
|
||||
) *compose.Lambda {
|
||||
return compose.InvokableLambda(func(ctx context.Context, sttOut STTOutput) ([]*schema.Message, error) {
|
||||
log := logger.Log
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
// 从 State 读取请求元数据
|
||||
state := stateFromCtx(ctx)
|
||||
@@ -48,7 +48,7 @@ func NewHistoryLambda(
|
||||
if userID != "" && scenarioRepo != nil {
|
||||
scenarios, err := scenarioRepo.FindByUserID(ctx, userID)
|
||||
if err != nil {
|
||||
log.Warnw("加载用户自建情景失败", "user_id", userID, "error", err)
|
||||
log.Warnw("load user scenarios failed", "user_id", userID, "error", err)
|
||||
} else if len(scenarios) > 0 {
|
||||
customScenarios = make(map[string]string, len(scenarios))
|
||||
customGreetings = make(map[string]string, len(scenarios))
|
||||
@@ -58,7 +58,7 @@ func NewHistoryLambda(
|
||||
customGreetings[s.ID] = s.Greeting
|
||||
}
|
||||
}
|
||||
log.Debugw("加载用户自建情景", "user_id", userID, "count", len(scenarios))
|
||||
log.Debugw("loaded user scenarios", "user_id", userID, "count", len(scenarios))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -78,7 +78,7 @@ func NewHistoryLambda(
|
||||
if historyFetcher != nil && sessionID != "" {
|
||||
history, err := historyFetcher(ctx, sessionID, maxHistory)
|
||||
if err != nil {
|
||||
log.Warnw("获取历史消息失败,继续处理", "error", err, "request_id", requestID)
|
||||
log.Warnw("fetch history failed, continuing", "error", err, "request_id", requestID)
|
||||
} else {
|
||||
for _, msg := range history {
|
||||
messages = append(messages, &schema.Message{
|
||||
@@ -121,8 +121,7 @@ func NewHistoryLambda(
|
||||
})
|
||||
}
|
||||
|
||||
log.Infow("历史组装完成",
|
||||
"request_id", requestID,
|
||||
log.Debugw("history assembled",
|
||||
"message_count", len(messages),
|
||||
"has_image", len(imageData) > 0,
|
||||
"scenario", scenario)
|
||||
|
||||
@@ -8,8 +8,9 @@ import (
|
||||
"github.com/cloudwego/eino/compose"
|
||||
|
||||
"github.com/hhs/camtalk/internal/ai/stt"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/hhs/camtalk/internal/util"
|
||||
)
|
||||
|
||||
// NewSTTLambda 创建 STT Lambda 节点。
|
||||
@@ -20,7 +21,7 @@ import (
|
||||
// 识别结果通过 Sender 发送 stt_result 到客户端。
|
||||
func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||
return compose.InvokableLambda(func(ctx context.Context, input PipelineInput) (STTOutput, error) {
|
||||
log := logger.Log
|
||||
log := trace.FromContext(ctx)
|
||||
sender := senderFromCtx(ctx)
|
||||
requestID := requestIDFromCtx(ctx)
|
||||
|
||||
@@ -39,8 +40,9 @@ func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||
|
||||
// 文本输入模式:跳过 STT
|
||||
if input.Text != "" {
|
||||
log.Infow("使用文本输入,跳过 STT",
|
||||
"request_id", requestID, "text", input.Text)
|
||||
log.Debugw("text input mode, skipping stt",
|
||||
"text_len", len(input.Text),
|
||||
"text_preview", util.Truncate(input.Text, 50))
|
||||
|
||||
// 发送 stt_result 保持前端消息流一致性
|
||||
if sender != nil {
|
||||
@@ -50,7 +52,7 @@ func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||
Text: input.Text,
|
||||
IsFinal: true,
|
||||
}); err != nil {
|
||||
log.Errorw("发送 stt_result 失败", "error", err)
|
||||
log.Errorw("send stt_result failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -73,8 +75,7 @@ func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||
return STTOutput{}, fmt.Errorf("stt: no audio data provided")
|
||||
}
|
||||
|
||||
log.Infow("开始语音识别",
|
||||
"request_id", requestID, "audio_bytes", len(input.AudioData))
|
||||
log.Debugw("stt recognition started", "audio_bytes", len(input.AudioData))
|
||||
|
||||
// 调用 STT 服务
|
||||
text, err := sttService.Recognize(ctx, input.AudioData, stt.Options{
|
||||
@@ -83,7 +84,7 @@ func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||
Language: input.Language,
|
||||
})
|
||||
if err != nil {
|
||||
log.Errorw("语音识别失败", "error", err, "request_id", requestID)
|
||||
log.Errorw("stt recognition failed", "error", err)
|
||||
if sender != nil {
|
||||
sender.SendError(models.WsError{
|
||||
Type: "error",
|
||||
@@ -97,11 +98,13 @@ func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||
|
||||
// STT 返回空文本
|
||||
if strings.TrimSpace(text) == "" {
|
||||
log.Infow("语音识别结果为空", "request_id", requestID)
|
||||
log.Infow("stt returned empty text")
|
||||
text = "(未识别到语音)"
|
||||
}
|
||||
|
||||
log.Infow("语音识别完成", "request_id", requestID, "text", text)
|
||||
log.Debugw("stt recognition completed",
|
||||
"text_len", len(text),
|
||||
"text_preview", util.Truncate(text, 50))
|
||||
|
||||
// 发送 stt_result
|
||||
if sender != nil {
|
||||
@@ -111,7 +114,7 @@ func NewSTTLambda(sttService stt.Service) *compose.Lambda {
|
||||
Text: text,
|
||||
IsFinal: true,
|
||||
}); err != nil {
|
||||
log.Errorw("发送 stt_result 失败", "error", err)
|
||||
log.Errorw("send stt_result failed", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -9,8 +9,8 @@ import (
|
||||
"github.com/cloudwego/eino/schema"
|
||||
|
||||
"github.com/hhs/camtalk/internal/ai/tts"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// NewTTSLambda 创建 TTS Transform Lambda 节点。
|
||||
@@ -26,7 +26,7 @@ func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, tts
|
||||
defer sw.Close()
|
||||
defer input.Close()
|
||||
|
||||
log := logger.Log
|
||||
log := trace.FromContext(ctx)
|
||||
sender := senderFromCtx(ctx)
|
||||
requestID := requestIDFromCtx(ctx)
|
||||
|
||||
@@ -48,7 +48,7 @@ func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, tts
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
log.Errorw("TTS: stream recv error", "error", err, "request_id", requestID)
|
||||
log.Errorw("TTS: stream recv error", "error", err)
|
||||
break
|
||||
}
|
||||
if sentence != "" {
|
||||
@@ -61,7 +61,7 @@ func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, tts
|
||||
return
|
||||
}
|
||||
|
||||
log.Infow("开始 TTS 合成", "request_id", requestID, "sentence_count", len(sentences))
|
||||
log.Infow("开始 TTS 合成", "sentence_count", len(sentences))
|
||||
|
||||
// 将句子数组转为 channel
|
||||
sentenceCh := make(chan string, len(sentences))
|
||||
@@ -78,7 +78,7 @@ func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, tts
|
||||
SampleRate: ttsSampleRate,
|
||||
})
|
||||
if err != nil {
|
||||
log.Errorw("TTS 合成启动失败(已跳过)", "error", err, "request_id", requestID)
|
||||
log.Errorw("TTS 合成启动失败(已跳过)", "error", err)
|
||||
sw.Send(struct{}{}, nil)
|
||||
return
|
||||
}
|
||||
@@ -87,7 +87,7 @@ func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, tts
|
||||
for chunk := range ttsStream {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
log.Infow("TTS 流被中断", "request_id", requestID)
|
||||
log.Debugw("tts stream interrupted")
|
||||
sw.Send(struct{}{}, ctx.Err())
|
||||
return
|
||||
default:
|
||||
@@ -107,7 +107,7 @@ func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, tts
|
||||
}
|
||||
}
|
||||
|
||||
log.Infow("TTS 合成完成", "request_id", requestID)
|
||||
log.Infow("TTS 合成完成")
|
||||
sw.Send(struct{}{}, nil)
|
||||
}()
|
||||
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// Middleware 返回 Gin 中间件,按 key 维度限流。
|
||||
@@ -26,6 +28,13 @@ func Middleware(limiter Limiter, keyFunc func(*gin.Context) string) gin.HandlerF
|
||||
allowed, retryAfter := limiter.Allow(c.Request.Context(), key)
|
||||
|
||||
if !allowed {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
log.Warnw("rate limited",
|
||||
"client_ip", c.ClientIP(),
|
||||
"path", c.Request.URL.Path,
|
||||
"limit_key", key,
|
||||
"retry_after_sec", int(retryAfter.Seconds()+0.5))
|
||||
|
||||
// 设置 Retry-After header(秒)
|
||||
c.Header("Retry-After", fmt.Sprintf("%d", int(retryAfter.Seconds()+0.5)))
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
@@ -72,6 +73,7 @@ func NewRedisLimiter(client *redis.Client, cfg config.RateLimitConfig) *RedisLim
|
||||
|
||||
// Allow 实现 Limiter 接口。
|
||||
func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||
log := trace.FromContext(ctx)
|
||||
cfg := l.getBucketConfig(key)
|
||||
|
||||
now := float64(time.Now().UnixNano()) / 1e9 // 秒,浮点
|
||||
@@ -81,6 +83,7 @@ func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Durati
|
||||
cfg.Capacity, cfg.Rate, now, ttl).Result()
|
||||
|
||||
if err != nil {
|
||||
log.Errorw("rate limit check failed", "key", key, "error", err)
|
||||
// Redis 错误时降级:允许请求(fail-open 策略)
|
||||
return true, 0
|
||||
}
|
||||
@@ -100,6 +103,7 @@ func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Durati
|
||||
}
|
||||
|
||||
retryAfter := time.Duration(retryAfterSec*1000) * time.Millisecond
|
||||
log.Warnw("rate limit triggered", "key", key, "retry_after_sec", retryAfterSec)
|
||||
return false, retryAfter
|
||||
}
|
||||
|
||||
|
||||
@@ -10,8 +10,9 @@ import (
|
||||
"github.com/google/uuid"
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
"github.com/hhs/camtalk/internal/util"
|
||||
)
|
||||
|
||||
// RedisManager 基于 Redis 的 SessionManager 实现。
|
||||
@@ -87,7 +88,8 @@ func (m *RedisManager) CreateWithID(ctx context.Context, id string, userID strin
|
||||
return "", fmt.Errorf("redis create session: %w", err)
|
||||
}
|
||||
|
||||
logger.Log.Debugw("redis session created", "session", id, "user_id", userID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Debugw("redis session created", "session_id", id, "user_id", userID)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
@@ -96,8 +98,11 @@ const placeholderHistoryMark = "__placeholder__"
|
||||
|
||||
// Get 获取会话。
|
||||
func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
vals, err := m.rdb.HGetAll(ctx, metaKey(sessionID)).Result()
|
||||
if err != nil {
|
||||
log.Errorw("redis get session failed", "session_id", sessionID, "error", err)
|
||||
return nil, fmt.Errorf("redis get session: %w", err)
|
||||
}
|
||||
if len(vals) == 0 {
|
||||
@@ -115,6 +120,7 @@ func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Sessi
|
||||
sess.Config.DetailLevel = vals["config.detail_level"]
|
||||
sess.Config.Language = vals["config.language"]
|
||||
|
||||
log.Debugw("redis session retrieved", "session_id", sessionID)
|
||||
return sess, nil
|
||||
}
|
||||
|
||||
@@ -150,7 +156,9 @@ func (m *RedisManager) UpdateConfig(ctx context.Context, sessionID string, patch
|
||||
|
||||
// 刷新 TTL
|
||||
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
logger.Log.Debugw("redis session config updated", "session", sessionID)
|
||||
|
||||
log := trace.FromContext(ctx)
|
||||
log.Debugw("redis session config updated", "session_id", sessionID)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -170,7 +178,9 @@ func (m *RedisManager) UpdateTitle(ctx context.Context, sessionID string, title
|
||||
}
|
||||
|
||||
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
logger.Log.Debugw("redis session title updated", "session", sessionID, "title", title)
|
||||
|
||||
log := trace.FromContext(ctx)
|
||||
log.Debugw("redis session title updated", "session_id", sessionID, "title", title)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -285,7 +295,11 @@ func (m *RedisManager) GetHistory(ctx context.Context, sessionID string, limit i
|
||||
}
|
||||
var msg models.Message
|
||||
if err := json.Unmarshal([]byte(raw), &msg); err != nil {
|
||||
logger.Log.Warnw("invalid history entry", "session", sessionID, "raw", raw)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("invalid history entry",
|
||||
"session_id", sessionID,
|
||||
"raw_len", len(raw),
|
||||
"raw_preview", util.Truncate(raw, 100))
|
||||
continue
|
||||
}
|
||||
msgs = append(msgs, msg)
|
||||
@@ -436,7 +450,8 @@ func (m *RedisManager) Destroy(ctx context.Context, sessionID string) error {
|
||||
m.rdb.SRem(ctx, userSessKey(userID), sessionID)
|
||||
}
|
||||
|
||||
logger.Log.Debugw("redis session destroyed", "session", sessionID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Debugw("redis session destroyed", "session_id", sessionID)
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// Redis key 前缀。
|
||||
@@ -83,7 +83,8 @@ func (r *CachedUserRepository) SaveRefreshToken(ctx context.Context, userID, tok
|
||||
pipe.Set(ctx, key, userID, ttl)
|
||||
pipe.SAdd(ctx, userRefreshKey(userID), tokenHash)
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
logger.Log.Warnw("Redis cache write failed for refresh token", "error", err)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache write failed for refresh token", "error", err)
|
||||
// 降级:DB 已写入成功,Redis 失败不影响正确性
|
||||
}
|
||||
return nil
|
||||
@@ -100,7 +101,8 @@ func (r *CachedUserRepository) FindRefreshToken(ctx context.Context, tokenHash s
|
||||
}
|
||||
// redis.Nil 表示 key 不存在,其他错误记录日志后降级到 DB
|
||||
if err != redis.Nil {
|
||||
logger.Log.Warnw("Redis cache read failed for refresh token", "error", err)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache read failed for refresh token", "error", err)
|
||||
}
|
||||
|
||||
// 降级到 DB
|
||||
@@ -139,7 +141,8 @@ func (r *CachedUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash
|
||||
pipe.SRem(ctx, userRefreshKey(userID), tokenHash)
|
||||
}
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
logger.Log.Warnw("Redis cache delete failed for refresh token", "error", err)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache delete failed for refresh token", "error", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -159,7 +162,8 @@ func (r *CachedUserRepository) DeleteUserRefreshTokens(ctx context.Context, user
|
||||
}
|
||||
keys = append(keys, userKey)
|
||||
if err := r.rdb.Del(ctx, keys...).Err(); err != nil {
|
||||
logger.Log.Warnw("Redis cache batch delete failed for user refresh tokens", "error", err, "userID", userID)
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache batch delete failed for user refresh tokens", "error", err, "user_id", userID)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// PgMessageRepository 基于 PostgreSQL 的 MessageRepository 实现。
|
||||
@@ -21,14 +22,24 @@ func NewPgMessageRepository(pool *pgxpool.Pool) *PgMessageRepository {
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) SaveMessage(ctx context.Context, sessionID string, msg models.Message, tokensUsed int) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
_, err := r.pool.Exec(ctx,
|
||||
`INSERT INTO messages (session_id, role, content, tokens_used) VALUES ($1, $2, $3, $4)`,
|
||||
sessionID, msg.Role, msg.Content, tokensUsed,
|
||||
)
|
||||
return err
|
||||
if err != nil {
|
||||
log.Errorw("save message failed", "session_id", sessionID, "role", msg.Role, "error", err)
|
||||
return err
|
||||
}
|
||||
|
||||
log.Debugw("message saved", "session_id", sessionID, "role", msg.Role, "tokens_used", tokensUsed)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string, limit int, beforeID int64) ([]StoredMessage, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
if limit <= 0 {
|
||||
limit = 50
|
||||
}
|
||||
@@ -56,6 +67,7 @@ func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string,
|
||||
)
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("get messages failed", "session_id", sessionID, "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -64,12 +76,16 @@ func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string,
|
||||
rows[i], rows[j] = rows[j], rows[i]
|
||||
}
|
||||
|
||||
log.Debugw("messages retrieved", "session_id", sessionID, "count", len(rows))
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) queryMessages(ctx context.Context, query string, args ...any) ([]StoredMessage, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
pgxRows, err := r.pool.Query(ctx, query, args...)
|
||||
if err != nil {
|
||||
log.Errorw("query messages failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
defer pgxRows.Close()
|
||||
@@ -78,17 +94,21 @@ func (r *PgMessageRepository) queryMessages(ctx context.Context, query string, a
|
||||
for pgxRows.Next() {
|
||||
var m StoredMessage
|
||||
if err := pgxRows.Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt); err != nil {
|
||||
log.Errorw("scan message row failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
messages = append(messages, m)
|
||||
}
|
||||
if err := pgxRows.Err(); err != nil {
|
||||
log.Errorw("iterate message rows failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
var m StoredMessage
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`SELECT id, session_id, role, content, tokens_used, created_at
|
||||
@@ -102,24 +122,34 @@ func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID stri
|
||||
return nil, ErrMessageNotFound
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("get last message failed", "session_id", sessionID, "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
log.Debugw("last message retrieved", "session_id", sessionID, "message_id", m.ID)
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) GetMessageCount(ctx context.Context, sessionID string) (int, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
var count int
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE session_id = $1`,
|
||||
sessionID,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
log.Errorw("get message count failed", "session_id", sessionID, "error", err)
|
||||
return 0, err
|
||||
}
|
||||
|
||||
log.Debugw("message count retrieved", "session_id", sessionID, "count", count)
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessionIDs []string) (map[string]SessionMessageStats, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
if len(sessionIDs) == 0 {
|
||||
return map[string]SessionMessageStats{}, nil
|
||||
}
|
||||
@@ -143,6 +173,7 @@ func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessio
|
||||
sessionIDs,
|
||||
)
|
||||
if err != nil {
|
||||
log.Errorw("get session message stats failed", "session_count", len(sessionIDs), "error", err)
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
@@ -152,12 +183,16 @@ func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessio
|
||||
var sid string
|
||||
var stats SessionMessageStats
|
||||
if err := rows.Scan(&sid, &stats.MessageCount, &stats.LastMessage); err != nil {
|
||||
log.Errorw("scan message stats row failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
result[sid] = stats
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
log.Errorw("iterate message stats rows failed", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
log.Debugw("session message stats retrieved", "session_count", len(sessionIDs), "result_count", len(result))
|
||||
return result, nil
|
||||
}
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// PgSessionRepository 基于 PostgreSQL 的 SessionRepository 实现。
|
||||
@@ -19,6 +21,8 @@ func NewPgSessionRepository(pool *pgxpool.Pool) *PgSessionRepository {
|
||||
}
|
||||
|
||||
func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
_, err := r.pool.Exec(ctx,
|
||||
`INSERT INTO sessions (id, user_id, title, config, created_at, updated_at)
|
||||
VALUES ($1, $2, $3, $4, $5, $6)
|
||||
@@ -28,10 +32,18 @@ func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
|
||||
updated_at = EXCLUDED.updated_at`,
|
||||
s.ID, s.UserID, s.Title, s.Config, s.CreatedAt, s.UpdatedAt,
|
||||
)
|
||||
return err
|
||||
if err != nil {
|
||||
log.Errorw("save session failed", "session_id", s.ID, "error", err)
|
||||
return err
|
||||
}
|
||||
|
||||
log.Debugw("session saved", "session_id", s.ID, "user_id", s.UserID)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PgSessionRepository) FindByID(ctx context.Context, id string) (*SessionRecord, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
var s SessionRecord
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`SELECT id, user_id, title, config, created_at, updated_at
|
||||
@@ -41,12 +53,17 @@ func (r *PgSessionRepository) FindByID(ctx context.Context, id string) (*Session
|
||||
return nil, ErrSessionNotFound
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find session failed", "session_id", id, "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
log.Debugw("session found", "session_id", id)
|
||||
return &s, nil
|
||||
}
|
||||
|
||||
func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, page, size int) ([]SessionRecord, int, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
if page <= 0 {
|
||||
page = 1
|
||||
}
|
||||
@@ -60,6 +77,7 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
|
||||
if err := r.pool.QueryRow(ctx,
|
||||
`SELECT COUNT(*) FROM sessions WHERE user_id = $1`, userID,
|
||||
).Scan(&total); err != nil {
|
||||
log.Errorw("count user sessions failed", "user_id", userID, "error", err)
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
@@ -73,6 +91,7 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
|
||||
userID, size, offset,
|
||||
)
|
||||
if err != nil {
|
||||
log.Errorw("find user sessions failed", "user_id", userID, "error", err)
|
||||
return nil, 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
@@ -81,66 +100,90 @@ func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, pag
|
||||
for rows.Next() {
|
||||
var s SessionRecord
|
||||
if err := rows.Scan(&s.ID, &s.UserID, &s.Title, &s.Config, &s.CreatedAt, &s.UpdatedAt); err != nil {
|
||||
log.Errorw("scan session row failed", "user_id", userID, "error", err)
|
||||
return nil, 0, err
|
||||
}
|
||||
list = append(list, s)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
log.Errorw("iterate session rows failed", "user_id", userID, "error", err)
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
log.Debugw("user sessions found", "user_id", userID, "count", len(list), "total", total)
|
||||
return list, total, nil
|
||||
}
|
||||
|
||||
func (r *PgSessionRepository) UpdateTitle(ctx context.Context, id string, title string) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
tag, err := r.pool.Exec(ctx,
|
||||
`UPDATE sessions SET title = $2, updated_at = NOW() WHERE id = $1`,
|
||||
id, title,
|
||||
)
|
||||
if err != nil {
|
||||
log.Errorw("update session title failed", "session_id", id, "error", err)
|
||||
return err
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
log.Debugw("session title updated", "session_id", id)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PgSessionRepository) UpdateConfig(ctx context.Context, id string, configJSON []byte) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
tag, err := r.pool.Exec(ctx,
|
||||
`UPDATE sessions SET config = $2, updated_at = NOW() WHERE id = $1`,
|
||||
id, configJSON,
|
||||
)
|
||||
if err != nil {
|
||||
log.Errorw("update session config failed", "session_id", id, "error", err)
|
||||
return err
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
log.Debugw("session config updated", "session_id", id)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PgSessionRepository) Touch(ctx context.Context, id string) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
tag, err := r.pool.Exec(ctx,
|
||||
`UPDATE sessions SET updated_at = NOW() WHERE id = $1`, id,
|
||||
)
|
||||
if err != nil {
|
||||
log.Errorw("touch session failed", "session_id", id, "error", err)
|
||||
return err
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
log.Debugw("session touched", "session_id", id)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PgSessionRepository) Delete(ctx context.Context, id string) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
tag, err := r.pool.Exec(ctx,
|
||||
`DELETE FROM sessions WHERE id = $1`, id,
|
||||
)
|
||||
if err != nil {
|
||||
log.Errorw("delete session failed", "session_id", id, "error", err)
|
||||
return err
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
log.Debugw("session deleted", "session_id", id)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -7,6 +7,8 @@ import (
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// PgUserRepository 基于 PostgreSQL 的 UserRepository 实现。
|
||||
@@ -20,18 +22,25 @@ func NewPgUserRepository(pool *pgxpool.Pool) *PgUserRepository {
|
||||
}
|
||||
|
||||
func (r *PgUserRepository) Create(ctx context.Context, username, passwordHash string) (string, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
var id string
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`INSERT INTO users (username, password_hash) VALUES ($1, $2) RETURNING id`,
|
||||
username, passwordHash,
|
||||
).Scan(&id)
|
||||
if err != nil {
|
||||
log.Errorw("create user failed", "username", username, "error", err)
|
||||
return "", err
|
||||
}
|
||||
|
||||
log.Debugw("user created", "user_id", id, "username", username)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func (r *PgUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
var u User
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`SELECT id, username, password_hash, created_at, updated_at FROM users WHERE username = $1`,
|
||||
@@ -41,12 +50,17 @@ func (r *PgUserRepository) FindByUsername(ctx context.Context, username string)
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find user by username failed", "username", username, "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
log.Debugw("user found by username", "user_id", u.ID, "username", username)
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
func (r *PgUserRepository) FindByID(ctx context.Context, id string) (*User, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
var u User
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`SELECT id, username, password_hash, created_at, updated_at FROM users WHERE id = $1`,
|
||||
@@ -56,20 +70,33 @@ func (r *PgUserRepository) FindByID(ctx context.Context, id string) (*User, erro
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find user by id failed", "user_id", id, "error", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
log.Debugw("user found by id", "user_id", id)
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
func (r *PgUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
_, err := r.pool.Exec(ctx,
|
||||
`INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)`,
|
||||
userID, tokenHash, expiresAt,
|
||||
)
|
||||
return err
|
||||
if err != nil {
|
||||
log.Errorw("save refresh token failed", "user_id", userID, "error", err)
|
||||
return err
|
||||
}
|
||||
|
||||
log.Debugw("refresh token saved", "user_id", userID)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PgUserRepository) FindRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
var userID string
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`SELECT user_id FROM refresh_tokens WHERE token_hash = $1 AND expires_at > NOW()`,
|
||||
@@ -79,23 +106,42 @@ func (r *PgUserRepository) FindRefreshToken(ctx context.Context, tokenHash strin
|
||||
return "", ErrRefreshTokenNotFound
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find refresh token failed", "error", err)
|
||||
return "", err
|
||||
}
|
||||
|
||||
log.Debugw("refresh token found", "user_id", userID)
|
||||
return userID, nil
|
||||
}
|
||||
|
||||
func (r *PgUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
_, err := r.pool.Exec(ctx,
|
||||
`DELETE FROM refresh_tokens WHERE token_hash = $1`,
|
||||
tokenHash,
|
||||
)
|
||||
return err
|
||||
if err != nil {
|
||||
log.Errorw("delete refresh token failed", "error", err)
|
||||
return err
|
||||
}
|
||||
|
||||
log.Debugw("refresh token deleted")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PgUserRepository) DeleteUserRefreshTokens(ctx context.Context, userID string) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
_, err := r.pool.Exec(ctx,
|
||||
`DELETE FROM refresh_tokens WHERE user_id = $1`,
|
||||
userID,
|
||||
)
|
||||
return err
|
||||
if err != nil {
|
||||
log.Errorw("delete user refresh tokens failed", "user_id", userID, "error", err)
|
||||
return err
|
||||
}
|
||||
|
||||
log.Debugw("user refresh tokens deleted", "user_id", userID)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// UserScenarioRepository 用户自建情景仓储接口。
|
||||
@@ -35,6 +36,8 @@ func NewPostgresUserScenarioRepo(pool *pgxpool.Pool) UserScenarioRepository {
|
||||
|
||||
// Create 创建用户情景。
|
||||
func (r *PostgresUserScenarioRepo) Create(ctx context.Context, scenario *models.UserScenario) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
INSERT INTO user_scenarios (id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at)
|
||||
VALUES ($1, $2, $3, $4, NULLIF($5, ''), $6, NULLIF($7, ''), $8, $9, $10)
|
||||
@@ -69,13 +72,18 @@ func (r *PostgresUserScenarioRepo) Create(ctx context.Context, scenario *models.
|
||||
).Scan(&scenario.ID, &scenario.CreatedAt, &scenario.UpdatedAt)
|
||||
|
||||
if err != nil {
|
||||
log.Errorw("create user scenario failed", "user_id", scenario.UserID, "name", scenario.Name, "error", err)
|
||||
return fmt.Errorf("create user scenario: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenario created", "scenario_id", scenario.ID, "user_id", scenario.UserID, "name", scenario.Name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// FindByID 根据 ID 查找情景。
|
||||
func (r *PostgresUserScenarioRepo) FindByID(ctx context.Context, id string) (*models.UserScenario, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||
FROM user_scenarios
|
||||
@@ -100,13 +108,18 @@ func (r *PostgresUserScenarioRepo) FindByID(ctx context.Context, id string) (*mo
|
||||
return nil, fmt.Errorf("user scenario not found: %s", id)
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find user scenario failed", "scenario_id", id, "error", err)
|
||||
return nil, fmt.Errorf("find user scenario: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenario found", "scenario_id", id)
|
||||
return &scenario, nil
|
||||
}
|
||||
|
||||
// FindByIDAndUserID 根据 ID 和用户 ID 查找情景(权限校验)。
|
||||
func (r *PostgresUserScenarioRepo) FindByIDAndUserID(ctx context.Context, id, userID string) (*models.UserScenario, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||
FROM user_scenarios
|
||||
@@ -131,13 +144,18 @@ func (r *PostgresUserScenarioRepo) FindByIDAndUserID(ctx context.Context, id, us
|
||||
return nil, fmt.Errorf("user scenario not found or no permission")
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find user scenario by id and user failed", "scenario_id", id, "user_id", userID, "error", err)
|
||||
return nil, fmt.Errorf("find user scenario: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenario found by id and user", "scenario_id", id, "user_id", userID)
|
||||
return &scenario, nil
|
||||
}
|
||||
|
||||
// FindByUserID 查找用户的所有情景。
|
||||
func (r *PostgresUserScenarioRepo) FindByUserID(ctx context.Context, userID string) ([]*models.UserScenario, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
SELECT id, user_id, name, icon, description, prompt, greeting, language, created_at, updated_at
|
||||
FROM user_scenarios
|
||||
@@ -147,6 +165,7 @@ func (r *PostgresUserScenarioRepo) FindByUserID(ctx context.Context, userID stri
|
||||
|
||||
rows, err := r.pool.Query(ctx, query, userID)
|
||||
if err != nil {
|
||||
log.Errorw("find user scenarios failed", "user_id", userID, "error", err)
|
||||
return nil, fmt.Errorf("find user scenarios: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
@@ -167,19 +186,25 @@ func (r *PostgresUserScenarioRepo) FindByUserID(ctx context.Context, userID stri
|
||||
&s.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
log.Errorw("scan user scenario row failed", "user_id", userID, "error", err)
|
||||
return nil, fmt.Errorf("scan user scenario: %w", err)
|
||||
}
|
||||
scenarios = append(scenarios, &s)
|
||||
}
|
||||
|
||||
if err = rows.Err(); err != nil {
|
||||
log.Errorw("iterate user scenarios failed", "user_id", userID, "error", err)
|
||||
return nil, fmt.Errorf("iterate user scenarios: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenarios found", "user_id", userID, "count", len(scenarios))
|
||||
return scenarios, nil
|
||||
}
|
||||
|
||||
// Update 更新用户情景。
|
||||
func (r *PostgresUserScenarioRepo) Update(ctx context.Context, scenario *models.UserScenario) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `
|
||||
UPDATE user_scenarios
|
||||
SET name = $1, icon = $2, description = $3, prompt = $4, greeting = $5, language = $6, updated_at = $7
|
||||
@@ -205,34 +230,47 @@ func (r *PostgresUserScenarioRepo) Update(ctx context.Context, scenario *models.
|
||||
return fmt.Errorf("user scenario not found or no permission")
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("update user scenario failed", "scenario_id", scenario.ID, "user_id", scenario.UserID, "error", err)
|
||||
return fmt.Errorf("update user scenario: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenario updated", "scenario_id", scenario.ID, "user_id", scenario.UserID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete 删除用户情景。
|
||||
func (r *PostgresUserScenarioRepo) Delete(ctx context.Context, id string) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `DELETE FROM user_scenarios WHERE id = $1`
|
||||
|
||||
result, err := r.pool.Exec(ctx, query, id)
|
||||
if err != nil {
|
||||
log.Errorw("delete user scenario failed", "scenario_id", id, "error", err)
|
||||
return fmt.Errorf("delete user scenario: %w", err)
|
||||
}
|
||||
|
||||
if result.RowsAffected() == 0 {
|
||||
return fmt.Errorf("user scenario not found")
|
||||
}
|
||||
|
||||
log.Debugw("user scenario deleted", "scenario_id", id)
|
||||
return nil
|
||||
}
|
||||
|
||||
// CountByUserID 统计用户的情景数量。
|
||||
func (r *PostgresUserScenarioRepo) CountByUserID(ctx context.Context, userID string) (int, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
query := `SELECT COUNT(*) FROM user_scenarios WHERE user_id = $1`
|
||||
|
||||
var count int
|
||||
err := r.pool.QueryRow(ctx, query, userID).Scan(&count)
|
||||
if err != nil {
|
||||
log.Errorw("count user scenarios failed", "user_id", userID, "error", err)
|
||||
return 0, fmt.Errorf("count user scenarios: %w", err)
|
||||
}
|
||||
|
||||
log.Debugw("user scenarios counted", "user_id", userID, "count", count)
|
||||
return count, nil
|
||||
}
|
||||
|
||||
46
backend/internal/trace/context.go
Normal file
46
backend/internal/trace/context.go
Normal file
@@ -0,0 +1,46 @@
|
||||
package trace
|
||||
|
||||
import "context"
|
||||
|
||||
type traceIDKey struct{}
|
||||
type requestIDKey struct{}
|
||||
type sessionIDKey struct{}
|
||||
|
||||
// WithTraceID 将 trace ID 注入 context(连接级/会话级标识)
|
||||
func WithTraceID(ctx context.Context, traceID string) context.Context {
|
||||
return context.WithValue(ctx, traceIDKey{}, traceID)
|
||||
}
|
||||
|
||||
// GetTraceID 从 context 提取 trace ID
|
||||
func GetTraceID(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(traceIDKey{}).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// WithRequestID 将 request ID 注入 context(单次请求/查询标识)
|
||||
func WithRequestID(ctx context.Context, requestID string) context.Context {
|
||||
return context.WithValue(ctx, requestIDKey{}, requestID)
|
||||
}
|
||||
|
||||
// GetRequestID 从 context 提取 request ID
|
||||
func GetRequestID(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(requestIDKey{}).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// WithSessionID 将 session ID 注入 context(会话存储标识)
|
||||
func WithSessionID(ctx context.Context, sessionID string) context.Context {
|
||||
return context.WithValue(ctx, sessionIDKey{}, sessionID)
|
||||
}
|
||||
|
||||
// GetSessionID 从 context 提取 session ID
|
||||
func GetSessionID(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(sessionIDKey{}).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
42
backend/internal/trace/eino_test.go
Normal file
42
backend/internal/trace/eino_test.go
Normal file
@@ -0,0 +1,42 @@
|
||||
package trace_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/cloudwego/eino/compose"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
func TestEinoContextPropagation(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
testTraceID := "01J5TEST123456789"
|
||||
ctx = trace.WithTraceID(ctx, testTraceID)
|
||||
|
||||
var capturedTraceID string
|
||||
|
||||
g := compose.NewGraph[string, string]()
|
||||
g.AddLambdaNode("test_node", compose.InvokableLambda(
|
||||
func(ctx context.Context, input string) (string, error) {
|
||||
capturedTraceID = trace.GetTraceID(ctx)
|
||||
return "ok", nil
|
||||
},
|
||||
))
|
||||
g.AddEdge(compose.START, "test_node")
|
||||
g.AddEdge("test_node", compose.END)
|
||||
|
||||
runnable, err := g.Compile(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("compile failed: %v", err)
|
||||
}
|
||||
|
||||
_, err = runnable.Invoke(ctx, "test_input")
|
||||
if err != nil {
|
||||
t.Fatalf("invoke failed: %v", err)
|
||||
}
|
||||
|
||||
if capturedTraceID != testTraceID {
|
||||
t.Errorf("trace_id lost in Eino propagation: got %q, want %q",
|
||||
capturedTraceID, testTraceID)
|
||||
}
|
||||
}
|
||||
63
backend/internal/trace/gin_logger.go
Normal file
63
backend/internal/trace/gin_logger.go
Normal file
@@ -0,0 +1,63 @@
|
||||
package trace
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GinLogger 记录每个 HTTP 请求的 method/path/status/latency
|
||||
func GinLogger() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
start := time.Now()
|
||||
path := c.Request.URL.Path
|
||||
query := c.Request.URL.RawQuery
|
||||
|
||||
c.Next()
|
||||
|
||||
latency := time.Since(start).Milliseconds()
|
||||
status := c.Writer.Status()
|
||||
log := FromContext(c.Request.Context())
|
||||
|
||||
fields := []interface{}{
|
||||
"method", c.Request.Method,
|
||||
"path", path,
|
||||
"status", status,
|
||||
"latency_ms", latency,
|
||||
"client_ip", c.ClientIP(),
|
||||
}
|
||||
if query != "" {
|
||||
fields = append(fields, "query", query)
|
||||
}
|
||||
if errStr := c.Errors.String(); errStr != "" {
|
||||
fields = append(fields, "errors", errStr)
|
||||
}
|
||||
|
||||
switch {
|
||||
case status >= 500:
|
||||
log.Errorw("request completed", fields...)
|
||||
case status >= 400:
|
||||
log.Warnw("request completed", fields...)
|
||||
default:
|
||||
log.Infow("request completed", fields...)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GinRecovery 自定义 panic 恢复中间件,使用 zap 记录
|
||||
func GinRecovery() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
log := FromContext(c.Request.Context())
|
||||
log.Errorw("panic recovered",
|
||||
"error", err,
|
||||
"path", c.Request.URL.Path,
|
||||
"method", c.Request.Method,
|
||||
"client_ip", c.ClientIP())
|
||||
c.AbortWithStatus(500)
|
||||
}
|
||||
}()
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
22
backend/internal/trace/id.go
Normal file
22
backend/internal/trace/id.go
Normal file
@@ -0,0 +1,22 @@
|
||||
package trace
|
||||
|
||||
import (
|
||||
cryptorand "crypto/rand"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/oklog/ulid/v2"
|
||||
)
|
||||
|
||||
var entropyPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
return ulid.Monotonic(cryptorand.Reader, 0)
|
||||
},
|
||||
}
|
||||
|
||||
// GenerateTraceID 生成并发安全的 ULID trace ID
|
||||
func GenerateTraceID() string {
|
||||
entropy := entropyPool.Get().(*ulid.MonotonicEntropy)
|
||||
defer entropyPool.Put(entropy)
|
||||
return ulid.MustNew(ulid.Timestamp(time.Now()), entropy).String()
|
||||
}
|
||||
25
backend/internal/trace/logger.go
Normal file
25
backend/internal/trace/logger.go
Normal file
@@ -0,0 +1,25 @@
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// FromContext 返回自动附加 trace_id/request_id/session_id 的 logger
|
||||
func FromContext(ctx context.Context) *zap.SugaredLogger {
|
||||
log := logger.Log
|
||||
|
||||
if traceID := GetTraceID(ctx); traceID != "" {
|
||||
log = log.With("trace_id", traceID)
|
||||
}
|
||||
if requestID := GetRequestID(ctx); requestID != "" {
|
||||
log = log.With("request_id", requestID)
|
||||
}
|
||||
if sessionID := GetSessionID(ctx); sessionID != "" {
|
||||
log = log.With("session_id", sessionID)
|
||||
}
|
||||
|
||||
return log
|
||||
}
|
||||
17
backend/internal/trace/middleware.go
Normal file
17
backend/internal/trace/middleware.go
Normal file
@@ -0,0 +1,17 @@
|
||||
package trace
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
// TraceMiddleware 为每个 HTTP 请求生成 trace ID 并注入 context
|
||||
func TraceMiddleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
traceID := GenerateTraceID()
|
||||
ctx := WithTraceID(c.Request.Context(), traceID)
|
||||
ctx = WithRequestID(ctx, traceID) // REST: trace_id == request_id
|
||||
|
||||
c.Request = c.Request.WithContext(ctx)
|
||||
c.Header("X-Trace-ID", traceID) // 返回给客户端用于排查
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
9
backend/internal/util/string.go
Normal file
9
backend/internal/util/string.go
Normal file
@@ -0,0 +1,9 @@
|
||||
package util
|
||||
|
||||
// Truncate 截断字符串到指定长度,超出部分用 "..." 替换
|
||||
func Truncate(s string, maxLen int) string {
|
||||
if len(s) <= maxLen {
|
||||
return s
|
||||
}
|
||||
return s[:maxLen] + "..."
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -8,11 +8,11 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"context"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"context"
|
||||
|
||||
"github.com/hhs/camtalk/internal/auth"
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
@@ -221,9 +221,9 @@ func TestWS_QueryFullFlow(t *testing.T) {
|
||||
imageB64 := base64.StdEncoding.EncodeToString([]byte("fake-image-data"))
|
||||
|
||||
mock := &MockOrchestrator{
|
||||
STTResult: "你好,世界",
|
||||
LLMDeltas: []string{"你好", ",世界!"},
|
||||
TTSAudios: []string{base64.StdEncoding.EncodeToString([]byte("mp3-data-1")), base64.StdEncoding.EncodeToString([]byte("mp3-data-2"))},
|
||||
STTResult: "你好,世界",
|
||||
LLMDeltas: []string{"你好", ",世界!"},
|
||||
TTSAudios: []string{base64.StdEncoding.EncodeToString([]byte("mp3-data-1")), base64.StdEncoding.EncodeToString([]byte("mp3-data-2"))},
|
||||
}
|
||||
|
||||
srv, wsURL := setupTestServer(t, mock)
|
||||
@@ -332,7 +332,7 @@ func TestWS_UnknownMessageType(t *testing.T) {
|
||||
err := conn.WriteJSON(map[string]string{"type": "unknown_type"})
|
||||
require.NoError(t, err)
|
||||
|
||||
errMsg := readJSON(t, conn)
|
||||
errMsg := readJSON(t, conn)
|
||||
assert.Equal(t, "error", errMsg["type"])
|
||||
assert.Equal(t, "INVALID_MESSAGE", errMsg["code"])
|
||||
assert.Contains(t, errMsg["message"], "unknown message type")
|
||||
@@ -642,7 +642,7 @@ func TestWS_AuthExpiredToken(t *testing.T) {
|
||||
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
||||
Session: config.SessionConfig{MaxHistory: 20},
|
||||
}
|
||||
r.GET("/ws", ServeWS(sessionMgr, &MockOrchestrator{}, cfg, tokenMgr, nil))
|
||||
r.GET("/ws", ServeWS(sessionMgr, &MockOrchestrator{}, cfg, tokenMgr, nil, nil))
|
||||
srv := httptest.NewServer(r)
|
||||
defer srv.Close()
|
||||
|
||||
|
||||
@@ -805,8 +805,9 @@ func (r *CachedUserRepository) DeleteRefreshToken(ctx, tokenHash) error {
|
||||
```
|
||||
|
||||
**降级策略**:
|
||||
- Redis 操作失败时记录日志,但不阻断主流程
|
||||
- Redis 操作失败时使用 `trace.FromContext(ctx)` 记录 Warn 日志(带 trace_id),但不阻断主流程
|
||||
- DB 是唯一真实数据源,Redis 仅用于加速
|
||||
- 详见 `docs/13-日志追踪.md` — 存储层日志实现
|
||||
|
||||
### 6. Gin 中间件实现
|
||||
|
||||
|
||||
@@ -360,7 +360,7 @@ if cfg.RateLimit.Enabled {
|
||||
logger.Log.Info("rate limiter initialized with Redis backend")
|
||||
} else {
|
||||
// 单实例:使用内存令牌桶
|
||||
limiter = ratelimit.NewLimiter(cfg.RateLimit)
|
||||
limiter = ratelimit.NewMemoryLimiter(cfg.RateLimit)
|
||||
logger.Log.Info("rate limiter initialized with in-memory backend")
|
||||
}
|
||||
defer limiter.Stop()
|
||||
@@ -590,9 +590,9 @@ case "query":
|
||||
// 限流检查
|
||||
if limiter != nil {
|
||||
key := fmt.Sprintf("%s:query", userID)
|
||||
allowed, retryAfter := limiter.Allow(context.Background(), key)
|
||||
allowed, retryAfter := limiter.Allow(ctx, key)
|
||||
if !allowed {
|
||||
logger.Log.Warnw("rate limited", "user_id", userID, "retry_after", retryAfter)
|
||||
// 限流触发时自动记录 Warn 日志(在 limiter 内部使用 trace.FromContext)
|
||||
errors.SendWSError(client, errors.CodeRateLimited, msg.RequestID,
|
||||
fmt.Errorf("rate limited, retry after %s", retryAfter.Round(time.Second)))
|
||||
continue
|
||||
@@ -605,7 +605,7 @@ case "query":
|
||||
**设计要点**:
|
||||
- key 格式:`userID:query`(用户级限流)
|
||||
- 拒绝时发送 `RATE_LIMITED` 错误到客户端
|
||||
- 记录警告日志(便于监控告警)
|
||||
- 限流触发时 `RedisLimiter.Allow` 内部自动记录 Warn 日志(带 trace_id,详见 `docs/13-日志追踪.md`)
|
||||
- 不阻塞其他消息类型(`ping`/`config`/`interrupt` 不限流)
|
||||
|
||||
### 配置加载与依赖注入
|
||||
|
||||
544
docs/13-日志追踪.md
Normal file
544
docs/13-日志追踪.md
Normal file
@@ -0,0 +1,544 @@
|
||||
# 日志追踪系统
|
||||
|
||||
## 概述
|
||||
|
||||
CamTalk 全链路日志追踪系统,通过统一的 trace ID 机制,将 REST API 和 WebSocket 两大入口的所有日志串联起来,实现分布式环境下的请求链路可观测性。
|
||||
|
||||
**核心目标**:
|
||||
- 统一 trace ID 贯穿 REST/WebSocket 两大入口
|
||||
- 所有日志自动附加 trace_id/request_id/session_id
|
||||
- 保护用户隐私,敏感文本截断或降级
|
||||
- 支持按 trace_id 快速定位完整请求链路
|
||||
|
||||
## Trace ID 作用域
|
||||
|
||||
| 标识 | 作用域 | 生成时机 | 用途 |
|
||||
|-----|--------|---------|------|
|
||||
| `trace_id` | **连接级**(整个 WebSocket 生命周期)<br/>**请求级**(单次 REST 请求) | REST: 中间件生成<br/>WebSocket: 升级时生成 | 关联同一连接/请求的所有日志 |
|
||||
| `session_id` | 会话级(对话上下文存储) | ServeWS 时生成 | 标识会话存储 |
|
||||
| `request_id` | 查询级(单次 WebSocket 查询) | 客户端每次查询传入 | 区分同一连接的不同查询 |
|
||||
|
||||
**WebSocket 场景示例**:用户打开页面建立 WebSocket,发起 3 次对话查询:
|
||||
|
||||
```
|
||||
连接建立 trace_id=01J5AAA session_id=uuid-123
|
||||
├─ 查询1 trace_id=01J5AAA request_id=req-001 (问天气)
|
||||
├─ 查询2 trace_id=01J5AAA request_id=req-002 (问新闻)
|
||||
└─ 查询3 trace_id=01J5AAA request_id=req-003 (问股票)
|
||||
```
|
||||
|
||||
**REST 场景示例**:
|
||||
|
||||
```
|
||||
POST /api/auth/login trace_id=01J5BBB request_id=01J5BBB
|
||||
GET /api/conversations trace_id=01J5CCC request_id=01J5CCC
|
||||
```
|
||||
|
||||
## 核心组件
|
||||
|
||||
```mermaid
|
||||
graph TB
|
||||
subgraph trace包["trace 包"]
|
||||
ID["id.go<br/>ULID 生成器"]
|
||||
CTX["context.go<br/>context key 管理"]
|
||||
LOG["logger.go<br/>context-aware logger"]
|
||||
MW["middleware.go<br/>Gin trace 中间件"]
|
||||
end
|
||||
|
||||
subgraph logger包["logger 包"]
|
||||
GINLOG["middleware.go<br/>Gin 请求日志"]
|
||||
GINREC["GinRecovery<br/>panic 恢复"]
|
||||
end
|
||||
|
||||
subgraph 入口层["入口层"]
|
||||
REST["REST API<br/>trace 中间件注入"]
|
||||
WS["WebSocket<br/>ServeWS 注入"]
|
||||
end
|
||||
|
||||
subgraph 业务层["业务层"]
|
||||
HANDLER["Handler"]
|
||||
ADAPTER["Eino Adapter"]
|
||||
NODES["Eino Nodes"]
|
||||
end
|
||||
|
||||
subgraph 存储层["存储层"]
|
||||
PG["PostgreSQL<br/>session/user/message/scenario"]
|
||||
REDIS["Redis<br/>session/cache/ratelimit"]
|
||||
end
|
||||
|
||||
ID --> MW
|
||||
CTX --> LOG
|
||||
LOG --> HANDLER
|
||||
LOG --> ADAPTER
|
||||
LOG --> NODES
|
||||
LOG --> PG
|
||||
LOG --> REDIS
|
||||
MW --> REST
|
||||
GINLOG --> REST
|
||||
WS --> LOG
|
||||
```
|
||||
|
||||
### trace/id.go — ULID 生成器
|
||||
|
||||
使用 ULID(Universally Unique Lexicographically Sortable Identifier)作为 trace ID:
|
||||
- 时间排序:前 48 位是毫秒时间戳,天然按时间排序
|
||||
- 唯一性:后 80 位随机数,冲突概率极低
|
||||
- 并发安全:使用 `crypto/rand` + `sync.Pool` 复用 entropy 对象
|
||||
|
||||
```go
|
||||
package trace
|
||||
|
||||
import (
|
||||
cryptorand "crypto/rand"
|
||||
"sync"
|
||||
"time"
|
||||
"github.com/oklog/ulid/v2"
|
||||
)
|
||||
|
||||
var entropyPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
return ulid.Monotonic(cryptorand.Reader, 0)
|
||||
},
|
||||
}
|
||||
|
||||
// GenerateTraceID 生成并发安全的 ULID trace ID
|
||||
func GenerateTraceID() string {
|
||||
entropy := entropyPool.Get().(*ulid.MonotonicEntropy)
|
||||
defer entropyPool.Put(entropy)
|
||||
return ulid.MustNew(ulid.Timestamp(time.Now()), entropy).String()
|
||||
}
|
||||
```
|
||||
|
||||
### trace/context.go — Context Key 管理
|
||||
|
||||
统一管理所有 trace 相关的 context key:
|
||||
|
||||
```go
|
||||
package trace
|
||||
|
||||
import "context"
|
||||
|
||||
type traceIDKey struct{}
|
||||
type requestIDKey struct{}
|
||||
type sessionIDKey struct{}
|
||||
|
||||
// WithTraceID 将 trace ID 注入 context
|
||||
func WithTraceID(ctx context.Context, traceID string) context.Context {
|
||||
return context.WithValue(ctx, traceIDKey{}, traceID)
|
||||
}
|
||||
|
||||
func GetTraceID(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(traceIDKey{}).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// 类似定义 WithRequestID/GetRequestID 和 WithSessionID/GetSessionID
|
||||
```
|
||||
|
||||
### trace/logger.go — Context-Aware Logger
|
||||
|
||||
自动从 context 提取 trace 字段并附加到日志:
|
||||
|
||||
```go
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// FromContext 返回自动附加 trace_id/request_id/session_id 的 logger
|
||||
func FromContext(ctx context.Context) *zap.SugaredLogger {
|
||||
log := logger.Log
|
||||
|
||||
if traceID := GetTraceID(ctx); traceID != "" {
|
||||
log = log.With("trace_id", traceID)
|
||||
}
|
||||
if requestID := GetRequestID(ctx); requestID != "" {
|
||||
log = log.With("request_id", requestID)
|
||||
}
|
||||
if sessionID := GetSessionID(ctx); sessionID != "" {
|
||||
log = log.With("session_id", sessionID)
|
||||
}
|
||||
|
||||
return log
|
||||
}
|
||||
```
|
||||
|
||||
**使用模式对比**:
|
||||
|
||||
```go
|
||||
// Before: 手动传递字段
|
||||
logger.Log.Infow("message", "session", sessionID, "request", requestID)
|
||||
|
||||
// After: 自动附加
|
||||
trace.FromContext(ctx).Infow("message")
|
||||
```
|
||||
|
||||
### trace/middleware.go — Gin Trace 中间件
|
||||
|
||||
为 REST 请求生成 trace ID 并注入 context:
|
||||
|
||||
```go
|
||||
package trace
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
// TraceMiddleware 为每个 HTTP 请求生成 trace ID 并注入 context
|
||||
func TraceMiddleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
traceID := GenerateTraceID()
|
||||
ctx := WithTraceID(c.Request.Context(), traceID)
|
||||
ctx = WithRequestID(ctx, traceID) // REST: trace_id == request_id
|
||||
|
||||
c.Request = c.Request.WithContext(ctx)
|
||||
c.Header("X-Trace-ID", traceID) // 返回给客户端用于排查
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### logger/middleware.go — 请求日志与 Panic 恢复
|
||||
|
||||
记录所有 HTTP 请求的 method/path/status/latency:
|
||||
|
||||
```go
|
||||
package logger
|
||||
|
||||
import (
|
||||
"time"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hhs/camtalk/internal/trace"
|
||||
)
|
||||
|
||||
// GinLogger 记录每个 HTTP 请求的基础信息
|
||||
func GinLogger() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
start := time.Now()
|
||||
path := c.Request.URL.Path
|
||||
|
||||
c.Next()
|
||||
|
||||
latency := time.Since(start).Milliseconds()
|
||||
status := c.Writer.Status()
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
|
||||
switch {
|
||||
case status >= 500:
|
||||
log.Errorw("request completed", "method", c.Request.Method,
|
||||
"path", path, "status", status, "latency_ms", latency)
|
||||
case status >= 400:
|
||||
log.Warnw("request completed", "method", c.Request.Method,
|
||||
"path", path, "status", status, "latency_ms", latency)
|
||||
default:
|
||||
log.Infow("request completed", "method", c.Request.Method,
|
||||
"path", path, "status", status, "latency_ms", latency)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GinRecovery 自定义 panic 恢复中间件
|
||||
func GinRecovery() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
log := trace.FromContext(c.Request.Context())
|
||||
log.Errorw("panic recovered", "error", err,
|
||||
"path", c.Request.URL.Path, "method", c.Request.Method)
|
||||
c.AbortWithStatus(500)
|
||||
}
|
||||
}()
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## 中间件注册顺序
|
||||
|
||||
在 `cmd/server/main.go` 中,三层中间件按顺序注册:
|
||||
|
||||
```go
|
||||
r := gin.New()
|
||||
r.Use(trace.TraceMiddleware()) // 第一层:生成 trace ID
|
||||
r.Use(logger.GinLogger()) // 第二层:记录请求
|
||||
r.Use(logger.GinRecovery()) // 第三层:panic 恢复
|
||||
```
|
||||
|
||||
## 日志输出示例
|
||||
|
||||
### REST 请求
|
||||
|
||||
```json
|
||||
{
|
||||
"level": "info",
|
||||
"ts": 1718956800.123,
|
||||
"msg": "login success",
|
||||
"trace_id": "01J5A2B3C4D5E6F7G8H9J0K1M",
|
||||
"request_id": "01J5A2B3C4D5E6F7G8H9J0K1M",
|
||||
"username": "test_user"
|
||||
}
|
||||
```
|
||||
|
||||
### WebSocket 查询链路(含存储层)
|
||||
|
||||
```json
|
||||
// 1. 查询接收
|
||||
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"query received"}
|
||||
|
||||
// 2. 会话加载(Redis)
|
||||
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"redis session retrieved", "session_id":"abc-123"}
|
||||
|
||||
// 3. STT 完成
|
||||
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"stt recognition completed", "text_len":45}
|
||||
|
||||
// 4. LLM 完成
|
||||
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"llm generation completed", "tokens":150}
|
||||
|
||||
// 5. 消息持久化(PostgreSQL)
|
||||
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"message saved", "role":"user", "tokens_used":45}
|
||||
{"level":"debug", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"message saved", "role":"assistant", "tokens_used":150}
|
||||
|
||||
// 6. Pipeline 完成
|
||||
{"level":"info", "trace_id":"01J5XXX", "session_id":"abc-123", "request_id":"req-456", "msg":"query processing completed", "latency_ms":2340}
|
||||
```
|
||||
|
||||
### 限流触发场景
|
||||
|
||||
```json
|
||||
{"level":"warn", "trace_id":"01J5YYY", "msg":"rate limit triggered", "key":"ratelimit:user-456:query", "retry_after_sec":2.5}
|
||||
```
|
||||
|
||||
## 日志查询操作
|
||||
|
||||
### 按 trace_id 查询完整链路
|
||||
|
||||
**本地开发(文件日志)**:
|
||||
```bash
|
||||
# 查看完整链路
|
||||
grep 'trace_id":"01J5XXX"' backend.log | jq .
|
||||
|
||||
# 查看链路时间线
|
||||
grep 'trace_id":"01J5XXX"' backend.log | jq -r '[.ts, .msg] | @tsv'
|
||||
```
|
||||
|
||||
**Grafana Loki**:
|
||||
```logql
|
||||
{app="camtalk-backend"}
|
||||
|= "trace_id=01J5XXX"
|
||||
| json
|
||||
| line_format "{{.ts}} [{{.level}}] {{.msg}}"
|
||||
```
|
||||
|
||||
### 查询慢请求(延迟 > 5s)
|
||||
|
||||
```logql
|
||||
{app="camtalk-backend"}
|
||||
| json
|
||||
| msg="query processing completed"
|
||||
| latency_ms > 5000
|
||||
```
|
||||
|
||||
### 查询数据库错误
|
||||
|
||||
```logql
|
||||
{app="camtalk-backend"}
|
||||
| json
|
||||
| level="error"
|
||||
| msg=~".*failed"
|
||||
| line_format "{{.trace_id}} {{.msg}} {{.error}}"
|
||||
```
|
||||
|
||||
### 查询 Redis 降级事件
|
||||
|
||||
```logql
|
||||
{app="camtalk-backend"}
|
||||
| json
|
||||
| level="warn"
|
||||
| msg=~"redis.*failed"
|
||||
```
|
||||
|
||||
### 查询错误率
|
||||
|
||||
```logql
|
||||
sum(count_over_time({app="camtalk-backend"} | json | level="error" [5m]))
|
||||
```
|
||||
|
||||
## 敏感内容处理规范
|
||||
|
||||
### 完全禁止记录
|
||||
|
||||
- 用户明文密码
|
||||
- JWT token 完整内容(仅记录 "token_present: true")
|
||||
- API Key 完整值(仅记录前 8 字符 + "...")
|
||||
|
||||
### 截断后记录(最多 50 字符)
|
||||
|
||||
- 用户输入文本 → `text_preview`
|
||||
- LLM 生成文本 → `text_preview`
|
||||
- STT 识别文本 → `text_preview`
|
||||
|
||||
**示例**:
|
||||
```go
|
||||
log.Debugw("stt recognition completed",
|
||||
"text_len", len(text),
|
||||
"text_preview", util.Truncate(text, 50))
|
||||
```
|
||||
|
||||
### 仅记录长度/大小
|
||||
|
||||
- 图片数据 → `image_size_bytes`
|
||||
- 音频数据 → `audio_size_bytes`
|
||||
|
||||
### 降级为 Debug 级别
|
||||
|
||||
所有包含用户文本预览的日志,生产环境默认不输出。
|
||||
|
||||
## 日志级别使用准则
|
||||
|
||||
| 场景 | 级别 | 示例 |
|
||||
|-----|------|-----|
|
||||
| 请求生命周期里程碑 | Info | `"query received"`, `"pipeline completed"` |
|
||||
| 中间步骤详情 | Debug | `"stt recognition completed"`, `"history assembled"` |
|
||||
| 敏感内容相关 | Debug | 所有包含用户文本的日志 |
|
||||
| 预期内的失败 | Warn | `"login failed"`, `"rate limited"` |
|
||||
| 系统错误 | Error | `"database query failed"`, `"tts synthesis failed"` |
|
||||
| 严重故障 | Error + stack | `"panic recovered"` |
|
||||
|
||||
## 存储层日志实现
|
||||
|
||||
### PostgreSQL Repository 层
|
||||
|
||||
所有数据库操作统一使用 `trace.FromContext(ctx)` 记录日志:
|
||||
|
||||
**已实现文件**:
|
||||
- `backend/internal/store/session_pg.go` — 会话 CRUD
|
||||
- `backend/internal/store/user_pg.go` — 用户与 refresh token 操作
|
||||
- `backend/internal/store/message_pg.go` — 对话消息存储
|
||||
- `backend/internal/store/user_scenario_repository.go` — 用户自定义情景
|
||||
|
||||
**日志策略**:
|
||||
|
||||
```go
|
||||
func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
_, err := r.pool.Exec(ctx, ...)
|
||||
if err != nil {
|
||||
log.Errorw("save session failed", "session_id", s.ID, "error", err)
|
||||
return err
|
||||
}
|
||||
|
||||
log.Debugw("session saved", "session_id", s.ID, "user_id", s.UserID)
|
||||
return nil
|
||||
}
|
||||
```
|
||||
|
||||
**NotFound 处理**:预期内的空结果不记录错误:
|
||||
|
||||
```go
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, ErrSessionNotFound // 不记录日志
|
||||
}
|
||||
if err != nil {
|
||||
log.Errorw("find session failed", "session_id", id, "error", err)
|
||||
return nil, err
|
||||
}
|
||||
```
|
||||
|
||||
### Redis 服务层
|
||||
|
||||
**已实现文件**:
|
||||
- `backend/internal/session/redis.go` — RedisManager(会话存储)
|
||||
- `backend/internal/store/cached_user.go` — CachedUserRepository(用户缓存装饰器)
|
||||
- `backend/internal/ratelimit/redis_bucket.go` — RedisLimiter(令牌桶限流器)
|
||||
|
||||
**会话存储日志**(`redis.go`):
|
||||
|
||||
```go
|
||||
func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
vals, err := m.rdb.HGetAll(ctx, metaKey(sessionID)).Result()
|
||||
if err != nil {
|
||||
log.Errorw("redis get session failed", "session_id", sessionID, "error", err)
|
||||
return nil, fmt.Errorf("redis get session: %w", err)
|
||||
}
|
||||
|
||||
if len(vals) == 0 {
|
||||
return nil, ErrSessionNotFound // 不记录日志
|
||||
}
|
||||
|
||||
log.Debugw("redis session retrieved", "session_id", sessionID)
|
||||
return session, nil
|
||||
}
|
||||
```
|
||||
|
||||
**缓存降级日志**(`cached_user.go`):
|
||||
|
||||
```go
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
log := trace.FromContext(ctx)
|
||||
log.Warnw("redis cache write failed for refresh token", "error", err)
|
||||
// 降级:DB 已写入成功,Redis 失败不影响正确性
|
||||
}
|
||||
```
|
||||
|
||||
**限流触发日志**(`redis_bucket.go`):
|
||||
|
||||
```go
|
||||
func (l *RedisLimiter) Allow(ctx context.Context, key string) (bool, time.Duration) {
|
||||
log := trace.FromContext(ctx)
|
||||
|
||||
result, err := l.script.Run(ctx, ...).Result()
|
||||
if err != nil {
|
||||
log.Errorw("rate limit check failed", "key", key, "error", err)
|
||||
return true, 0 // fail-open 策略
|
||||
}
|
||||
|
||||
if allowed == 0 {
|
||||
log.Warnw("rate limit triggered", "key", key, "retry_after_sec", retryAfterSec)
|
||||
return false, retryAfter
|
||||
}
|
||||
|
||||
return true, 0
|
||||
}
|
||||
```
|
||||
|
||||
**级别选择原则**:
|
||||
- **Error**:Redis 连接失败、Lua 脚本执行失败(影响功能)
|
||||
- **Warn**:缓存写入失败(可降级)、限流触发(预期内异常)
|
||||
- **Debug**:正常操作完成(避免 Info 级别噪音)
|
||||
|
||||
## 编码规范
|
||||
|
||||
1. **日志语言**:统一使用英文
|
||||
2. **结构化**:始终使用 `Infow`/`Errorw`/`Warnw`/`Debugw`
|
||||
3. **Context 传递**:使用 `trace.FromContext(ctx)` 而非直接引用 `logger.Log`
|
||||
4. **敏感内容**:禁止在 Info 及以上级别记录用户文本原文
|
||||
5. **错误日志**:采用"调用方记录"原则,底层函数 return wrapped error
|
||||
6. **级别约定**:
|
||||
- `Debug`:内部状态跟踪、开发调试信息(数据库/缓存成功操作)
|
||||
- `Info`:请求/连接生命周期、关键操作里程碑
|
||||
- `Warn`:可降级异常(Redis 故障、限流触发)
|
||||
- `Error`:影响用户的操作失败(数据库错误、Redis 连接失败)
|
||||
- `Fatal`:仅启动阶段不可恢复错误
|
||||
7. **预期内的空结果**:`pgx.ErrNoRows`、`redis.Nil` 等不记录错误日志
|
||||
|
||||
## 性能考量
|
||||
|
||||
### FromContext 开销
|
||||
|
||||
- 有 trace_id:~200-300 ns/op
|
||||
- 无 trace_id:~10-20 ns/op(仅返回全局 logger)
|
||||
- 1000 QPS 场景额外开销约 0.2ms,可接受
|
||||
|
||||
### ULID 生成吞吐量
|
||||
|
||||
- 单线程:~500k ops/s
|
||||
- 并发 8 线程:~2M ops/s
|
||||
|
||||
**验收标准**:1000 QPS 下,trace 系统开销 < 1% CPU,< 0.5ms P99 延迟。
|
||||
@@ -17,6 +17,8 @@ CamTalk 是一款多模态实时 AI 视觉对话助手。用户通过摄像头
|
||||
| [09-情景切换](09-情景切换.md) | 多情景 AI 角色扮演系统(面试官、英语老师、辩论对手、翻译员、自由对话) |
|
||||
| [10-鉴权体系](10-鉴权体系.md) | JWT 双 token 轮转认证、bcrypt 密码哈希、Refresh Token Rotation、安全机制 |
|
||||
| [11-令牌桶限流](11-令牌桶限流.md) | 令牌桶限流算法、内存/Redis 双实现、Gin 中间件、WebSocket query 限流 |
|
||||
| [12-自定义情景](12-自定义情景.md) | 用户自定义情景的完整设计 |
|
||||
| [13-日志追踪](13-日志追踪.md) | 全链路日志追踪系统(trace ID、敏感内容保护、日志规范) |
|
||||
|
||||
|
||||
## 推荐阅读顺序
|
||||
@@ -30,6 +32,8 @@ CamTalk 是一款多模态实时 AI 视觉对话助手。用户通过摄像头
|
||||
7. **09-情景切换** — 多情景 AI 角色扮演系统
|
||||
8. **10-鉴权体系** — 认证授权机制详细设计
|
||||
9. **11-令牌桶限流** — 速率限制设计
|
||||
10. **12-自定义情景** — 用户自定义情景
|
||||
11. **13-日志追踪** — 全链路日志追踪(trace ID、敏感内容保护、开发参考)
|
||||
|
||||
## 功能扩展方向
|
||||
|
||||
|
||||
Reference in New Issue
Block a user