package store import ( "context" "errors" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" "github.com/hhs/camtalk/internal/models" "github.com/hhs/camtalk/internal/trace" ) // PgMessageRepository 基于 PostgreSQL 的 MessageRepository 实现。 type PgMessageRepository struct { pool *pgxpool.Pool } // NewPgMessageRepository 创建 PgMessageRepository。 func NewPgMessageRepository(pool *pgxpool.Pool) *PgMessageRepository { return &PgMessageRepository{pool: pool} } 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, ) 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 } var rows []StoredMessage var err error if beforeID > 0 { rows, err = r.queryMessages(ctx, `SELECT id, session_id, role, content, tokens_used, created_at FROM messages WHERE session_id = $1 AND id < $2 ORDER BY id DESC LIMIT $3`, sessionID, beforeID, limit, ) } else { rows, err = r.queryMessages(ctx, `SELECT id, session_id, role, content, tokens_used, created_at FROM messages WHERE session_id = $1 ORDER BY id DESC LIMIT $2`, sessionID, limit, ) } if err != nil { log.Errorw("get messages failed", "session_id", sessionID, "error", err) return nil, err } // 反转为升序 for i, j := 0, len(rows)-1; i < j; i, j = i+1, j-1 { 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() messages := make([]StoredMessage, 0) 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 FROM messages WHERE session_id = $1 ORDER BY id DESC LIMIT 1`, sessionID, ).Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt) if errors.Is(err, pgx.ErrNoRows) { 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 } rows, err := r.pool.Query(ctx, `WITH stats AS ( SELECT session_id, COUNT(*) AS cnt FROM messages WHERE session_id = ANY($1) GROUP BY session_id ), last_msg AS ( SELECT DISTINCT ON (session_id) session_id, content FROM messages WHERE session_id = ANY($1) ORDER BY session_id, id DESC ) SELECT s.session_id, s.cnt, COALESCE(lm.content, '') FROM stats s LEFT JOIN last_msg lm ON lm.session_id = s.session_id`, sessionIDs, ) if err != nil { log.Errorw("get session message stats failed", "session_count", len(sessionIDs), "error", err) return nil, err } defer rows.Close() result := make(map[string]SessionMessageStats) for rows.Next() { 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 }