diff --git a/backend/internal/store/message_pg.go b/backend/internal/store/message_pg.go new file mode 100644 index 0000000..8c1512a --- /dev/null +++ b/backend/internal/store/message_pg.go @@ -0,0 +1,120 @@ +package store + +import ( + "context" + "errors" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/hhs/camtalk/internal/models" +) + +// 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 { + _, 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 +} + +func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string, limit int, beforeID int64) ([]StoredMessage, error) { + 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 { + 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] + } + + return rows, nil +} + +func (r *PgMessageRepository) queryMessages(ctx context.Context, query string, args ...any) ([]StoredMessage, error) { + pgxRows, err := r.pool.Query(ctx, query, args...) + if err != nil { + return nil, err + } + defer pgxRows.Close() + + var messages []StoredMessage + for pgxRows.Next() { + var m StoredMessage + if err := pgxRows.Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt); err != nil { + return nil, err + } + messages = append(messages, m) + } + if err := pgxRows.Err(); err != nil { + return nil, err + } + return messages, nil +} + +func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error) { + 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 { + return nil, err + } + return &m, nil +} + +func (r *PgMessageRepository) GetMessageCount(ctx context.Context, sessionID string) (int, error) { + var count int + err := r.pool.QueryRow(ctx, + `SELECT COUNT(*) FROM messages WHERE session_id = $1`, + sessionID, + ).Scan(&count) + if err != nil { + return 0, err + } + return count, nil +}