Compare commits

...

10 Commits

Author SHA1 Message Date
hhs
239f8f9877 docs: 同步限流和鉴权文档的日志实现说明
- 更新限流文档:limiter 内部使用 trace.FromContext 自动记录日志
- 更新鉴权文档:Redis 降级策略使用 trace-aware 日志
- 引用 13-日志追踪.md 作为详细说明
- 移除过时的手动 logger.Log 调用示例
2026-06-21 23:19:11 +08:00
hhs
d430e6e5b2 docs: 补充存储层和限流器日志实现说明
添加 PostgreSQL、Redis、限流器三个模块的日志实现文档:
- PostgreSQL 4 个 repository 的日志策略和代码示例
- Redis 会话存储、缓存装饰器、限流器的日志级别选择
- 存储层日志查询示例(数据库错误、Redis 降级)
- 更新架构图,添加存储层节点
2026-06-21 23:11:17 +08:00
hhs
edc66625ba feat: 为 redis_bucket.go 添加限流日志
添加 trace-aware 日志:
- Error: Redis 限流检查失败(fail-open 降级)
- Warn: 限流触发,记录 key 和 retry_after_sec
2026-06-21 23:08:00 +08:00
hhs
9dce107a84 feat: 升级 cached_user.go 日志为 trace-aware
替换 logger.Log 为 trace.FromContext(ctx):
- SaveRefreshToken: Redis 缓存写入失败降级日志
- FindRefreshToken: Redis 缓存读取失败降级日志
- DeleteRefreshToken: Redis 缓存删除失败降级日志
- DeleteUserRefreshTokens: Redis 批量删除失败降级日志
2026-06-21 23:07:49 +08:00
hhs
1c7dd708a0 feat: 升级 session/redis.go 日志为 trace-aware
替换 logger.Log 为 trace.FromContext(ctx):
- CreateWithID: session 创建日志
- Get: session 获取日志(新增错误日志)
- UpdateConfig: 配置更新日志
- UpdateTitle: 标题更新日志
- GetHistory: 无效历史条目警告日志
- Destroy: session 销毁日志
2026-06-21 23:07:02 +08:00
hhs
d0e4bdaeec feat: 为 PostgreSQL store 层添加 trace 日志
为 4 个 PostgreSQL repository 添加 trace-aware 日志:
- session_pg.go: Save/Find/Update/Delete 操作日志
- user_pg.go: 用户 CRUD 和 refresh token 管理日志
- message_pg.go: 消息存储和查询日志
- user_scenario_repository.go: 自定义情景 CRUD 日志

日志策略:
- Error: 数据库操作失败
- Debug: 操作成功(避免 Info 级别噪音)
- NotFound (ErrNoRows) 不记录错误日志
2026-06-21 23:06:50 +08:00
hhs
55f7f183a7 docs: 完善日志追踪文档 2026-06-21 22:51:14 +08:00
hhs
6b4b033df3 feat: 日志级别优化(Phase 7)
- nodes_stt.go: STT 识别开始降为 Debug
- nodes_history.go: 历史组装完成降为 Debug
- nodes_tts.go: TTS 流中断降为 Debug
- 保持关键里程碑为 Info:query completed、tts synthesis started/completed
- 中间步骤详情降为 Debug,减少生产环境日志噪音
2026-06-21 22:37:23 +08:00
hhs
76d331c885 feat: Eino nodes 迁移到 trace 包(Phase 6.2)
- nodes_stt.go 使用 trace.FromContext 替换 logger.Log
- nodes_history.go 使用 trace.FromContext
- nodes_tts.go 使用 trace.FromContext
- nodes_done.go 使用 trace.FromContext
- 移除所有 nodes 中的 request_id 手动字段(自动附加)
- 所有日志消息改为英文
2026-06-21 22:36:15 +08:00
hhs
ad700743ef feat: Eino adapter 和 callback 迁移到 trace 包(Phase 6.1)
- 移除 adapter.go 中的 ctxKeySessionID 定义
- 移除 callback.go 中的 ctxKeyRequestID 定义
- 统一使用 trace.WithSessionID/WithRequestID
- adapter.go 使用 trace.FromContext 替换 logger.Log
- callback.go 使用 trace.FromContext
- 移除双重日志,SetActiveRequest 失败直接返回错误
- 更新测试文件导入 trace 包
2026-06-21 22:33:49 +08:00
19 changed files with 802 additions and 90 deletions

View File

@@ -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. 调用 GraphStream 模式 + 运行时 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
}

View File

@@ -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)

View File

@@ -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))

View File

@@ -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,

View File

@@ -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)

View File

@@ -8,8 +8,8 @@ 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"
)
@@ -21,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)
@@ -41,7 +41,6 @@ func NewSTTLambda(sttService stt.Service) *compose.Lambda {
// 文本输入模式:跳过 STT
if input.Text != "" {
log.Debugw("text input mode, skipping stt",
"request_id", requestID,
"text_len", len(input.Text),
"text_preview", util.Truncate(input.Text, 50))
@@ -53,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)
}
}
@@ -76,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{
@@ -86,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",
@@ -100,12 +98,11 @@ 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.Debugw("stt recognition completed",
"request_id", requestID,
"text_len", len(text),
"text_preview", util.Truncate(text, 50))
@@ -117,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)
}
}

View File

@@ -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)
}()

View File

@@ -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
}

View File

@@ -10,8 +10,8 @@ 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"
)
@@ -88,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
}
@@ -97,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 {
@@ -116,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
}
@@ -151,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
}
@@ -171,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
}
@@ -286,8 +295,9 @@ 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,
log := trace.FromContext(ctx)
log.Warnw("invalid history entry",
"session_id", sessionID,
"raw_len", len(raw),
"raw_preview", util.Truncate(raw, 100))
continue
@@ -440,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
}

View File

@@ -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)
}
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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()

View File

@@ -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 中间件实现

View File

@@ -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
View 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 生成器
使用 ULIDUniversally 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 延迟。

View File

@@ -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、敏感内容保护、开发参考
## 功能扩展方向