Files
CamTalk/backend/internal/store/user_pg.go

148 lines
3.8 KiB
Go
Raw Normal View History

2026-06-14 16:55:12 +08:00
package store
import (
"context"
"errors"
"time"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/trace"
2026-06-14 16:55:12 +08:00
)
// PgUserRepository 基于 PostgreSQL 的 UserRepository 实现。
type PgUserRepository struct {
pool *pgxpool.Pool
}
// NewPgUserRepository 创建 PgUserRepository。
func NewPgUserRepository(pool *pgxpool.Pool) *PgUserRepository {
return &PgUserRepository{pool: pool}
}
func (r *PgUserRepository) Create(ctx context.Context, username, passwordHash string) (string, error) {
log := trace.FromContext(ctx)
2026-06-14 16:55:12 +08:00
var id string
err := r.pool.QueryRow(ctx,
`INSERT INTO users (username, password_hash) VALUES ($1, $2) RETURNING id`,
username, passwordHash,
).Scan(&id)
if err != nil {
log.Errorw("create user failed", "username", username, "error", err)
2026-06-14 16:55:12 +08:00
return "", err
}
log.Debugw("user created", "user_id", id, "username", username)
2026-06-14 16:55:12 +08:00
return id, nil
}
func (r *PgUserRepository) FindByUsername(ctx context.Context, username string) (*User, error) {
log := trace.FromContext(ctx)
2026-06-14 16:55:12 +08:00
var u User
err := r.pool.QueryRow(ctx,
`SELECT id, username, password_hash, created_at, updated_at FROM users WHERE username = $1`,
username,
).Scan(&u.ID, &u.Username, &u.PasswordHash, &u.CreatedAt, &u.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
return nil, ErrUserNotFound
}
if err != nil {
log.Errorw("find user by username failed", "username", username, "error", err)
2026-06-14 16:55:12 +08:00
return nil, err
}
log.Debugw("user found by username", "user_id", u.ID, "username", username)
2026-06-14 16:55:12 +08:00
return &u, nil
}
func (r *PgUserRepository) FindByID(ctx context.Context, id string) (*User, error) {
log := trace.FromContext(ctx)
2026-06-14 16:55:12 +08:00
var u User
err := r.pool.QueryRow(ctx,
`SELECT id, username, password_hash, created_at, updated_at FROM users WHERE id = $1`,
id,
).Scan(&u.ID, &u.Username, &u.PasswordHash, &u.CreatedAt, &u.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
return nil, ErrUserNotFound
}
if err != nil {
log.Errorw("find user by id failed", "user_id", id, "error", err)
2026-06-14 16:55:12 +08:00
return nil, err
}
log.Debugw("user found by id", "user_id", id)
2026-06-14 16:55:12 +08:00
return &u, nil
}
func (r *PgUserRepository) SaveRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
log := trace.FromContext(ctx)
2026-06-14 16:55:12 +08:00
_, err := r.pool.Exec(ctx,
`INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)`,
userID, tokenHash, expiresAt,
)
if err != nil {
log.Errorw("save refresh token failed", "user_id", userID, "error", err)
return err
}
log.Debugw("refresh token saved", "user_id", userID)
return nil
2026-06-14 16:55:12 +08:00
}
func (r *PgUserRepository) FindRefreshToken(ctx context.Context, tokenHash string) (string, error) {
log := trace.FromContext(ctx)
2026-06-14 16:55:12 +08:00
var userID string
err := r.pool.QueryRow(ctx,
`SELECT user_id FROM refresh_tokens WHERE token_hash = $1 AND expires_at > NOW()`,
tokenHash,
).Scan(&userID)
if errors.Is(err, pgx.ErrNoRows) {
return "", ErrRefreshTokenNotFound
}
if err != nil {
log.Errorw("find refresh token failed", "error", err)
2026-06-14 16:55:12 +08:00
return "", err
}
log.Debugw("refresh token found", "user_id", userID)
2026-06-14 16:55:12 +08:00
return userID, nil
}
func (r *PgUserRepository) DeleteRefreshToken(ctx context.Context, tokenHash string) error {
log := trace.FromContext(ctx)
2026-06-14 16:55:12 +08:00
_, err := r.pool.Exec(ctx,
`DELETE FROM refresh_tokens WHERE token_hash = $1`,
tokenHash,
)
if err != nil {
log.Errorw("delete refresh token failed", "error", err)
return err
}
log.Debugw("refresh token deleted")
return nil
2026-06-14 16:55:12 +08:00
}
func (r *PgUserRepository) DeleteUserRefreshTokens(ctx context.Context, userID string) error {
log := trace.FromContext(ctx)
2026-06-14 16:55:12 +08:00
_, err := r.pool.Exec(ctx,
`DELETE FROM refresh_tokens WHERE user_id = $1`,
userID,
)
if err != nil {
log.Errorw("delete user refresh tokens failed", "user_id", userID, "error", err)
return err
}
log.Debugw("user refresh tokens deleted", "user_id", userID)
return nil
2026-06-14 16:55:12 +08:00
}