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) 不记录错误日志
This commit is contained in:
hhs
2026-06-21 23:06:50 +08:00
parent 55f7f183a7
commit d0e4bdaeec
4 changed files with 167 additions and 5 deletions

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
}