2026-06-14 17:53:14 +08:00
|
|
|
package store
|
|
|
|
|
|
|
|
|
|
import (
|
|
|
|
|
"context"
|
|
|
|
|
"errors"
|
|
|
|
|
|
|
|
|
|
"github.com/jackc/pgx/v5"
|
|
|
|
|
"github.com/jackc/pgx/v5/pgxpool"
|
|
|
|
|
|
|
|
|
|
"github.com/hhs/camtalk/internal/models"
|
2026-06-21 23:06:50 +08:00
|
|
|
"github.com/hhs/camtalk/internal/trace"
|
2026-06-14 17:53:14 +08:00
|
|
|
)
|
|
|
|
|
|
|
|
|
|
// 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 {
|
2026-06-21 23:06:50 +08:00
|
|
|
log := trace.FromContext(ctx)
|
|
|
|
|
|
2026-06-14 17:53:14 +08:00
|
|
|
_, 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,
|
|
|
|
|
)
|
2026-06-21 23:06:50 +08:00
|
|
|
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
|
2026-06-14 17:53:14 +08:00
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string, limit int, beforeID int64) ([]StoredMessage, error) {
|
2026-06-21 23:06:50 +08:00
|
|
|
log := trace.FromContext(ctx)
|
|
|
|
|
|
2026-06-14 17:53:14 +08:00
|
|
|
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 {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("get messages failed", "session_id", sessionID, "error", err)
|
2026-06-14 17:53:14 +08:00
|
|
|
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]
|
|
|
|
|
}
|
|
|
|
|
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Debugw("messages retrieved", "session_id", sessionID, "count", len(rows))
|
2026-06-14 17:53:14 +08:00
|
|
|
return rows, nil
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (r *PgMessageRepository) queryMessages(ctx context.Context, query string, args ...any) ([]StoredMessage, error) {
|
2026-06-21 23:06:50 +08:00
|
|
|
log := trace.FromContext(ctx)
|
|
|
|
|
|
2026-06-14 17:53:14 +08:00
|
|
|
pgxRows, err := r.pool.Query(ctx, query, args...)
|
|
|
|
|
if err != nil {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("query messages failed", "error", err)
|
2026-06-14 17:53:14 +08:00
|
|
|
return nil, err
|
|
|
|
|
}
|
|
|
|
|
defer pgxRows.Close()
|
|
|
|
|
|
2026-06-22 12:39:37 +08:00
|
|
|
messages := make([]StoredMessage, 0)
|
2026-06-14 17:53:14 +08:00
|
|
|
for pgxRows.Next() {
|
|
|
|
|
var m StoredMessage
|
|
|
|
|
if err := pgxRows.Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt); err != nil {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("scan message row failed", "error", err)
|
2026-06-14 17:53:14 +08:00
|
|
|
return nil, err
|
|
|
|
|
}
|
|
|
|
|
messages = append(messages, m)
|
|
|
|
|
}
|
|
|
|
|
if err := pgxRows.Err(); err != nil {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("iterate message rows failed", "error", err)
|
2026-06-14 17:53:14 +08:00
|
|
|
return nil, err
|
|
|
|
|
}
|
|
|
|
|
return messages, nil
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error) {
|
2026-06-21 23:06:50 +08:00
|
|
|
log := trace.FromContext(ctx)
|
|
|
|
|
|
2026-06-14 17:53:14 +08:00
|
|
|
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 {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("get last message failed", "session_id", sessionID, "error", err)
|
2026-06-14 17:53:14 +08:00
|
|
|
return nil, err
|
|
|
|
|
}
|
2026-06-21 23:06:50 +08:00
|
|
|
|
|
|
|
|
log.Debugw("last message retrieved", "session_id", sessionID, "message_id", m.ID)
|
2026-06-14 17:53:14 +08:00
|
|
|
return &m, nil
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (r *PgMessageRepository) GetMessageCount(ctx context.Context, sessionID string) (int, error) {
|
2026-06-21 23:06:50 +08:00
|
|
|
log := trace.FromContext(ctx)
|
|
|
|
|
|
2026-06-14 17:53:14 +08:00
|
|
|
var count int
|
|
|
|
|
err := r.pool.QueryRow(ctx,
|
|
|
|
|
`SELECT COUNT(*) FROM messages WHERE session_id = $1`,
|
|
|
|
|
sessionID,
|
|
|
|
|
).Scan(&count)
|
|
|
|
|
if err != nil {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("get message count failed", "session_id", sessionID, "error", err)
|
2026-06-14 17:53:14 +08:00
|
|
|
return 0, err
|
|
|
|
|
}
|
2026-06-21 23:06:50 +08:00
|
|
|
|
|
|
|
|
log.Debugw("message count retrieved", "session_id", sessionID, "count", count)
|
2026-06-14 17:53:14 +08:00
|
|
|
return count, nil
|
|
|
|
|
}
|
2026-06-14 17:58:42 +08:00
|
|
|
|
|
|
|
|
func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessionIDs []string) (map[string]SessionMessageStats, error) {
|
2026-06-21 23:06:50 +08:00
|
|
|
log := trace.FromContext(ctx)
|
|
|
|
|
|
2026-06-14 17:58:42 +08:00
|
|
|
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 {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("get session message stats failed", "session_count", len(sessionIDs), "error", err)
|
2026-06-14 17:58:42 +08:00
|
|
|
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 {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("scan message stats row failed", "error", err)
|
2026-06-14 17:58:42 +08:00
|
|
|
return nil, err
|
|
|
|
|
}
|
|
|
|
|
result[sid] = stats
|
|
|
|
|
}
|
|
|
|
|
if err := rows.Err(); err != nil {
|
2026-06-21 23:06:50 +08:00
|
|
|
log.Errorw("iterate message stats rows failed", "error", err)
|
2026-06-14 17:58:42 +08:00
|
|
|
return nil, err
|
|
|
|
|
}
|
2026-06-21 23:06:50 +08:00
|
|
|
|
|
|
|
|
log.Debugw("session message stats retrieved", "session_count", len(sessionIDs), "result_count", len(result))
|
2026-06-14 17:58:42 +08:00
|
|
|
return result, nil
|
|
|
|
|
}
|