Merge pull request 'feat: 添加 PostgreSQL 服务并挂载数据库迁移脚本' #101
120
backend/internal/store/message_pg.go
Normal file
120
backend/internal/store/message_pg.go
Normal file
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user