feat: added db transactions abstraction; cleaned up user registration; repos adapted new abstraction: all repos now support transactions, if present in ctx
This commit is contained in:
@@ -11,17 +11,13 @@ import (
|
||||
)
|
||||
|
||||
type AuthRepo struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewAuthRepo(db *sql.DB) *AuthRepo {
|
||||
return &AuthRepo{db: db}
|
||||
Repo
|
||||
}
|
||||
|
||||
func (r *AuthRepo) GetUserById(ctx context.Context, id string) (*domain.AuthUser, error) {
|
||||
user := &domain.AuthUser{}
|
||||
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"SELECT id, email, name, password_hash FROM users WHERE id = $1",
|
||||
id,
|
||||
).Scan(&user.ID, &user.Email, &user.Name, &user.PasswordHash)
|
||||
@@ -35,7 +31,7 @@ func (r *AuthRepo) GetUserById(ctx context.Context, id string) (*domain.AuthUser
|
||||
func (r *AuthRepo) GetUserByEmail(ctx context.Context, email string) (*domain.AuthUser, error) {
|
||||
user := &domain.AuthUser{}
|
||||
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"SELECT id, email, name, password_hash FROM users WHERE email = $1",
|
||||
email,
|
||||
).Scan(&user.ID, &user.Email, &user.Name, &user.PasswordHash)
|
||||
@@ -47,7 +43,7 @@ func (r *AuthRepo) GetUserByEmail(ctx context.Context, email string) (*domain.Au
|
||||
}
|
||||
|
||||
func (r *AuthRepo) StoreRefreshToken(ctx context.Context, userID string, tokenHash string, expiresAt time.Time) error {
|
||||
_, err := r.db.ExecContext(ctx,
|
||||
_, err := r.conn(ctx).ExecContext(ctx,
|
||||
"INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)",
|
||||
userID, tokenHash, expiresAt,
|
||||
)
|
||||
@@ -60,7 +56,7 @@ func (r *AuthRepo) StoreRefreshToken(ctx context.Context, userID string, tokenHa
|
||||
|
||||
func (r *AuthRepo) ConsumeRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||
var userID string
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"DELETE FROM refresh_tokens WHERE token_hash = $1 AND expires_at > NOW() RETURNING user_id",
|
||||
tokenHash,
|
||||
).Scan(&userID)
|
||||
@@ -75,7 +71,7 @@ func (r *AuthRepo) ConsumeRefreshToken(ctx context.Context, tokenHash string) (s
|
||||
}
|
||||
|
||||
func (r *AuthRepo) RevokeRefreshToken(ctx context.Context, tokenHash string) error {
|
||||
_, err := r.db.ExecContext(ctx, "DELETE FROM refresh_tokens WHERE token_hash = $1", tokenHash)
|
||||
_, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM refresh_tokens WHERE token_hash = $1", tokenHash)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to revoke refresh token: %w", err)
|
||||
}
|
||||
@@ -84,7 +80,7 @@ func (r *AuthRepo) RevokeRefreshToken(ctx context.Context, tokenHash string) err
|
||||
}
|
||||
|
||||
func (r *AuthRepo) RevokeRefreshTokens(ctx context.Context, userID string) error {
|
||||
_, err := r.db.ExecContext(ctx, "DELETE FROM refresh_tokens WHERE user_id = $1", userID)
|
||||
_, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM refresh_tokens WHERE user_id = $1", userID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to revoke refresh tokens: %w", err)
|
||||
}
|
||||
|
||||
@@ -11,16 +11,12 @@ import (
|
||||
)
|
||||
|
||||
type InviteRepo struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewInviteRepo(db *sql.DB) *InviteRepo {
|
||||
return &InviteRepo{db: db}
|
||||
Repo
|
||||
}
|
||||
|
||||
func (r *InviteRepo) CreateInvite(ctx context.Context, inviterUserID string, code string, expiresAt time.Time) (*domain.Invite, error) {
|
||||
var id string
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"INSERT INTO invites (inviter_user_id, code, expires_at) VALUES ($1, $2, $3) RETURNING id",
|
||||
inviterUserID, code, expiresAt,
|
||||
).Scan(&id)
|
||||
@@ -37,7 +33,7 @@ func (r *InviteRepo) CreateInvite(ctx context.Context, inviterUserID string, cod
|
||||
}
|
||||
|
||||
func (r *InviteRepo) DeleteInvite(ctx context.Context, userID string, inviteID string) error {
|
||||
_, err := r.db.ExecContext(ctx,
|
||||
_, err := r.conn(ctx).ExecContext(ctx,
|
||||
"DELETE FROM invites WHERE id = $1 AND inviter_user_id = $2 AND consumed_at IS NULL",
|
||||
inviteID, userID,
|
||||
)
|
||||
@@ -49,7 +45,7 @@ func (r *InviteRepo) DeleteInvite(ctx context.Context, userID string, inviteID s
|
||||
}
|
||||
|
||||
func (r *InviteRepo) ConsumeInvite(ctx context.Context, inviteID string, inviteeUserID string) error {
|
||||
res, err := r.db.ExecContext(ctx,
|
||||
res, err := r.conn(ctx).ExecContext(ctx,
|
||||
"UPDATE invites SET invitee_user_id=$1, consumed_at=NOW() WHERE id=$2 AND expires_at > NOW() AND consumed_at IS NULL",
|
||||
inviteeUserID, inviteID,
|
||||
)
|
||||
@@ -70,7 +66,7 @@ func (r *InviteRepo) ConsumeInvite(ctx context.Context, inviteID string, invitee
|
||||
func (r *InviteRepo) GetInvite(ctx context.Context, code string) (*domain.Invite, error) {
|
||||
var invite domain.Invite
|
||||
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"SELECT id, code, expires_at, consumed_at FROM invites WHERE code=$1",
|
||||
code,
|
||||
).Scan(&invite.ID, &invite.Code, &invite.ExpiresAt, &invite.ConsumedAt)
|
||||
@@ -82,7 +78,7 @@ func (r *InviteRepo) GetInvite(ctx context.Context, code string) (*domain.Invite
|
||||
}
|
||||
|
||||
func (r *InviteRepo) GetInvites(ctx context.Context, userID string, offset int, count int) ([]domain.Invite, error) {
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
rows, err := r.conn(ctx).QueryContext(ctx,
|
||||
"SELECT id, code, expires_at, consumed_at FROM invites WHERE inviter_user_id=$1 ORDER BY created_at DESC OFFSET $2 LIMIT $3",
|
||||
userID, offset, count,
|
||||
)
|
||||
@@ -110,7 +106,7 @@ func (r *InviteRepo) GetInvites(ctx context.Context, userID string, offset int,
|
||||
|
||||
func (r *InviteRepo) CountInvites(ctx context.Context, userID string) (int, error) {
|
||||
var count int
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"SELECT COUNT(*) FROM invites WHERE inviter_user_id=$1",
|
||||
userID,
|
||||
).Scan(&count)
|
||||
|
||||
+14
-46
@@ -5,29 +5,18 @@ import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"github.com/robindittmar/dttmr-api/internal/domain"
|
||||
)
|
||||
|
||||
type ListRepo struct {
|
||||
db *sql.DB
|
||||
Repo
|
||||
}
|
||||
|
||||
func NewListRepo(db *sql.DB) *ListRepo {
|
||||
return &ListRepo{db: db}
|
||||
}
|
||||
|
||||
func (r *ListRepo) CreateList(ctx context.Context, name string, userIDs []string) (*domain.List, error) {
|
||||
tx, err := r.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("begin transaction: %w", err)
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
func (r *ListRepo) CreateList(ctx context.Context, name string) (*domain.List, error) {
|
||||
list := &domain.List{Name: name}
|
||||
|
||||
err = tx.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"INSERT INTO lists (name) VALUES ($1) RETURNING id, created_at, modified_at",
|
||||
name,
|
||||
).Scan(&list.ID, &list.CreatedAt, &list.ModifiedAt)
|
||||
@@ -35,32 +24,11 @@ func (r *ListRepo) CreateList(ctx context.Context, name string, userIDs []string
|
||||
return nil, fmt.Errorf("failed to insert list: %w", err)
|
||||
}
|
||||
|
||||
stmt, err := tx.PrepareContext(ctx, "INSERT INTO list_users (list_id, user_id) VALUES ($1, $2)")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to prepare user/list association statement: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
err := stmt.Close()
|
||||
if err != nil {
|
||||
slog.Error("failed to close user/list association statement", slog.Any("error", err))
|
||||
}
|
||||
}()
|
||||
|
||||
for _, userID := range userIDs {
|
||||
if _, err = stmt.ExecContext(ctx, list.ID, userID); err != nil {
|
||||
return nil, fmt.Errorf("failed to insert user/list association: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, fmt.Errorf("commit transaction: %w", err)
|
||||
}
|
||||
|
||||
return list, nil
|
||||
}
|
||||
|
||||
func (r *ListRepo) DeleteList(ctx context.Context, listID string) error {
|
||||
_, err := r.db.ExecContext(ctx, "DELETE FROM lists WHERE id = $1", listID)
|
||||
_, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM lists WHERE id = $1", listID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete list: %w", err)
|
||||
}
|
||||
@@ -69,7 +37,7 @@ func (r *ListRepo) DeleteList(ctx context.Context, listID string) error {
|
||||
}
|
||||
|
||||
func (r *ListRepo) GetLists(ctx context.Context, userID string) ([]domain.List, error) {
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
rows, err := r.conn(ctx).QueryContext(ctx,
|
||||
"SELECT l.id, l.name, l.created_at, l.modified_at, (SELECT COUNT(*) FROM list_items WHERE list_id=l.id), (SELECT COUNT(*) FROM list_items WHERE list_id=l.id AND is_completed=true) FROM lists AS l INNER JOIN list_users ON l.id=list_users.list_id WHERE list_users.user_id = $1",
|
||||
userID,
|
||||
)
|
||||
@@ -96,7 +64,7 @@ func (r *ListRepo) GetLists(ctx context.Context, userID string) ([]domain.List,
|
||||
}
|
||||
|
||||
func (r *ListRepo) AddUserToList(ctx context.Context, listID string, userID string) error {
|
||||
_, err := r.db.ExecContext(ctx,
|
||||
_, err := r.conn(ctx).ExecContext(ctx,
|
||||
"INSERT INTO list_users (list_id, user_id) VALUES ($1, $2)",
|
||||
listID, userID,
|
||||
)
|
||||
@@ -108,7 +76,7 @@ func (r *ListRepo) AddUserToList(ctx context.Context, listID string, userID stri
|
||||
}
|
||||
|
||||
func (r *ListRepo) RemoveUserFromList(ctx context.Context, listID string, userID string) error {
|
||||
_, err := r.db.ExecContext(ctx,
|
||||
_, err := r.conn(ctx).ExecContext(ctx,
|
||||
"DELETE FROM list_users WHERE list_id = $1 AND user_id = $2",
|
||||
listID, userID,
|
||||
)
|
||||
@@ -122,7 +90,7 @@ func (r *ListRepo) RemoveUserFromList(ctx context.Context, listID string, userID
|
||||
func (r *ListRepo) IsUserInList(ctx context.Context, listID string, userID string) (bool, error) {
|
||||
var cnt int
|
||||
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"SELECT COUNT(*) FROM list_users WHERE list_id = $1 AND user_id = $2",
|
||||
listID, userID,
|
||||
).Scan(&cnt)
|
||||
@@ -136,7 +104,7 @@ func (r *ListRepo) IsUserInList(ctx context.Context, listID string, userID strin
|
||||
func (r *ListRepo) IsUserInListByItemID(ctx context.Context, listItemID string, userID string) (bool, error) {
|
||||
var cnt int
|
||||
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"SELECT COUNT(*) FROM list_users WHERE list_id = (SELECT list_id FROM list_items WHERE id = $1) AND user_id = $2",
|
||||
listItemID, userID,
|
||||
).Scan(&cnt)
|
||||
@@ -150,7 +118,7 @@ func (r *ListRepo) IsUserInListByItemID(ctx context.Context, listItemID string,
|
||||
func (r *ListRepo) CreateListItem(ctx context.Context, listID string, title string) (*domain.ListItem, error) {
|
||||
l := &domain.ListItem{Title: title}
|
||||
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"INSERT INTO list_items (list_id, title) VALUES ($1, $2) RETURNING id, is_completed, created_at, modified_at",
|
||||
listID, title,
|
||||
).Scan(&l.ID, &l.IsCompleted, &l.CreatedAt, &l.ModifiedAt)
|
||||
@@ -162,7 +130,7 @@ func (r *ListRepo) CreateListItem(ctx context.Context, listID string, title stri
|
||||
}
|
||||
|
||||
func (r *ListRepo) DeleteListItem(ctx context.Context, listItemID string) error {
|
||||
_, err := r.db.ExecContext(ctx, "DELETE FROM list_items WHERE id = $1", listItemID)
|
||||
_, err := r.conn(ctx).ExecContext(ctx, "DELETE FROM list_items WHERE id = $1", listItemID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete list item: %w", err)
|
||||
}
|
||||
@@ -171,7 +139,7 @@ func (r *ListRepo) DeleteListItem(ctx context.Context, listItemID string) error
|
||||
}
|
||||
|
||||
func (r *ListRepo) UpdateListItem(ctx context.Context, listItemID string, title string, isCompleted bool) error {
|
||||
_, err := r.db.ExecContext(ctx, "UPDATE list_items SET title = $1, is_completed = $2, modified_at = NOW() WHERE id = $3",
|
||||
_, err := r.conn(ctx).ExecContext(ctx, "UPDATE list_items SET title = $1, is_completed = $2, modified_at = NOW() WHERE id = $3",
|
||||
title, isCompleted, listItemID,
|
||||
)
|
||||
if err != nil {
|
||||
@@ -182,7 +150,7 @@ func (r *ListRepo) UpdateListItem(ctx context.Context, listItemID string, title
|
||||
}
|
||||
|
||||
func (r *ListRepo) SetListItemCompleted(ctx context.Context, listItemID string, isCompleted bool) error {
|
||||
_, err := r.db.ExecContext(ctx, "UPDATE list_items SET is_completed = $1, modified_at = NOW() WHERE id = $2",
|
||||
_, err := r.conn(ctx).ExecContext(ctx, "UPDATE list_items SET is_completed = $1, modified_at = NOW() WHERE id = $2",
|
||||
isCompleted, listItemID,
|
||||
)
|
||||
if err != nil {
|
||||
@@ -193,7 +161,7 @@ func (r *ListRepo) SetListItemCompleted(ctx context.Context, listItemID string,
|
||||
}
|
||||
|
||||
func (r *ListRepo) GetListItems(ctx context.Context, listID string) ([]domain.ListItem, error) {
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
rows, err := r.conn(ctx).QueryContext(ctx,
|
||||
"SELECT id, title, is_completed, created_at, modified_at FROM list_items WHERE list_id = $1 ORDER BY is_completed, modified_at DESC",
|
||||
listID,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
)
|
||||
|
||||
type Transaction interface {
|
||||
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
|
||||
QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
|
||||
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
|
||||
PrepareContext(ctx context.Context, query string) (*sql.Stmt, error)
|
||||
}
|
||||
|
||||
type Repo struct {
|
||||
*Transactor
|
||||
}
|
||||
|
||||
func NewRepo(t *Transactor) Repo {
|
||||
return Repo{t}
|
||||
}
|
||||
|
||||
func (r *Repo) conn(ctx context.Context) Transaction {
|
||||
if tx, ok := ctx.Value(txKey{}).(*sql.Tx); ok {
|
||||
return tx
|
||||
}
|
||||
return r.db
|
||||
}
|
||||
|
||||
type txKey struct{}
|
||||
|
||||
type Transactor struct {
|
||||
db *sql.DB
|
||||
sp atomic.Uint64
|
||||
}
|
||||
|
||||
func NewTransactor(db *sql.DB) *Transactor {
|
||||
return &Transactor{db: db, sp: atomic.Uint64{}}
|
||||
}
|
||||
|
||||
func (t *Transactor) WithinTx(ctx context.Context, fn func(context.Context) error) error {
|
||||
if tx, ok := ctx.Value(txKey{}).(*sql.Tx); ok {
|
||||
return t.withinSavepoint(ctx, tx, fn)
|
||||
}
|
||||
|
||||
tx, err := t.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
if err = fn(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
func (t *Transactor) withinSavepoint(ctx context.Context, tx *sql.Tx, fn func(context.Context) error) error {
|
||||
name := fmt.Sprintf("sp_%d", t.sp.Add(1))
|
||||
if _, err := tx.ExecContext(ctx, "SAVEPOINT "+name); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := fn(ctx); err != nil {
|
||||
_, _ = tx.ExecContext(ctx, "ROLLBACK TO SAVEPOINT"+name)
|
||||
return err
|
||||
}
|
||||
|
||||
_, err := tx.ExecContext(ctx, "RELEASE SAVEPOINT"+name)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package repository
|
||||
|
||||
import "database/sql"
|
||||
|
||||
type Store struct {
|
||||
*Transactor
|
||||
Auth *AuthRepo
|
||||
Invite *InviteRepo
|
||||
List *ListRepo
|
||||
User *UserRepo
|
||||
}
|
||||
|
||||
func NewStore(db *sql.DB) *Store {
|
||||
t := NewTransactor(db)
|
||||
r := NewRepo(t)
|
||||
return &Store{
|
||||
Transactor: t,
|
||||
Auth: &AuthRepo{r},
|
||||
Invite: &InviteRepo{r},
|
||||
List: &ListRepo{r},
|
||||
User: &UserRepo{r},
|
||||
}
|
||||
}
|
||||
@@ -2,30 +2,19 @@ package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"github.com/robindittmar/dttmr-api/internal/domain"
|
||||
)
|
||||
|
||||
type UserRepo struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewUserRepo(db *sql.DB) *UserRepo {
|
||||
return &UserRepo{db: db}
|
||||
Repo
|
||||
}
|
||||
|
||||
func (r *UserRepo) CreateUser(ctx context.Context, email string, name string, passwordHash string) (*domain.User, error) {
|
||||
tx, err := r.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("begin transaction: %w", err)
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
user := &domain.User{Email: email, Name: name}
|
||||
|
||||
err = tx.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"INSERT INTO users (email, name, password_hash) VALUES ($1, $2, $3) RETURNING id, created_at",
|
||||
email, name, passwordHash,
|
||||
).Scan(&user.ID, &user.CreatedAt)
|
||||
@@ -33,15 +22,11 @@ func (r *UserRepo) CreateUser(ctx context.Context, email string, name string, pa
|
||||
return nil, fmt.Errorf("failed to insert user: %w", err)
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, fmt.Errorf("commit transaction: %w", err)
|
||||
}
|
||||
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func (r *UserRepo) DeleteUser(ctx context.Context, userID string) error {
|
||||
_, err := r.db.ExecContext(ctx,
|
||||
_, err := r.conn(ctx).ExecContext(ctx,
|
||||
"DELETE FROM users WHERE id = $1",
|
||||
userID,
|
||||
)
|
||||
@@ -53,7 +38,7 @@ func (r *UserRepo) DeleteUser(ctx context.Context, userID string) error {
|
||||
}
|
||||
|
||||
func (r *UserRepo) ChangePassword(ctx context.Context, userID string, passwordHash string) error {
|
||||
_, err := r.db.ExecContext(ctx,
|
||||
_, err := r.conn(ctx).ExecContext(ctx,
|
||||
"UPDATE users SET password_hash = $1 WHERE id = $2",
|
||||
passwordHash, userID,
|
||||
)
|
||||
@@ -67,7 +52,7 @@ func (r *UserRepo) ChangePassword(ctx context.Context, userID string, passwordHa
|
||||
func (r *UserRepo) GetUserByEmail(ctx context.Context, email string) (*domain.User, error) {
|
||||
user := &domain.User{}
|
||||
|
||||
err := r.db.QueryRowContext(ctx,
|
||||
err := r.conn(ctx).QueryRowContext(ctx,
|
||||
"SELECT id, email, name FROM users WHERE email = $1",
|
||||
email,
|
||||
).Scan(&user.ID, &user.Email, &user.Name)
|
||||
|
||||
Reference in New Issue
Block a user