uow setup

This commit is contained in:
Kristian Borgwarth 2026-03-26 22:41:34 +01:00
parent 85e91699be
commit 79066ee638
5 changed files with 36 additions and 37 deletions

View file

@ -2,6 +2,7 @@ package handlers
import ( import (
"bytes" "bytes"
"context"
"encoding/json" "encoding/json"
"os" "os"
@ -18,11 +19,12 @@ type createNoteCommand struct {
} }
type CreateNoteHandler struct { type CreateNoteHandler struct {
uow repositories.UnitOfWork
noteRepo repositories.NoteRepository noteRepo repositories.NoteRepository
tagRepo repositories.TagRepository tagRepo repositories.TagRepository
} }
func (h CreateNoteHandler) Handle(params []byte) (*rpc.Response, *rpc.Error) { func (h CreateNoteHandler) Handle(ctx context.Context, params []byte) (*rpc.Response, *rpc.Error) {
var cmd createNoteCommand var cmd createNoteCommand
if err := json.Unmarshal(params, &cmd); err != nil { if err := json.Unmarshal(params, &cmd); err != nil {

View file

View file

@ -1,21 +1,22 @@
package repositories package repositories
import "database/sql" import (
"context"
)
type NoteRepository interface { type NoteRepository interface {
Upsert(title string, path string, slug string) error Upsert(ctx context.Context, title string, path string, slug string) error
} }
type noteRepository struct { type noteRepository struct {
db *sql.DB dbCtx DBContext
} }
func NewNoteRepository(db *sql.DB) NoteRepository { func NewNoteRepository(dbContext DBContext) NoteRepository {
return &noteRepository{db: db} return &noteRepository{dbCtx: dbContext}
} }
func (r *noteRepository) Upsert(ctx context.Context, title string, path string, slug string) error {
func (r *noteRepository) Upsert(title string, path string, slug string) error {
query := ` query := `
INSERT INTO notes (title, path, slug) INSERT INTO notes (title, path, slug)
VALUES ($1, $2, $3) VALUES ($1, $2, $3)
@ -23,6 +24,6 @@ func (r *noteRepository) Upsert(title string, path string, slug string) error {
SET title = EXCLUDED.title, SET title = EXCLUDED.title,
path = EXCLUDED.path; path = EXCLUDED.path;
` `
_, err := r.db.Exec(query, title, path, slug) _, err := r.dbCtx.ExecContext(ctx, query, title, path, slug)
return err return err
} }

View file

@ -1,33 +1,27 @@
package repositories package repositories
import ( import (
"database/sql" "context"
"strings" "strings"
) )
type TagRepository interface { type TagRepository interface {
Upsert(names []string) error Upsert(ctx context.Context, names []string) error
} }
type tagRepository struct { type tagRepository struct {
db *sql.DB dbContext DBContext
} }
func NewTagRepository(db *sql.DB) TagRepository { func NewTagRepository(dbContext DBContext) TagRepository {
return &tagRepository{db: db} return &tagRepository{dbContext: dbContext}
} }
func (r *tagRepository) Upsert(names []string) error { func (r *tagRepository) Upsert(ctx context.Context, names []string) error {
if len(names) == 0 { if len(names) == 0 {
return nil return nil
} }
tx, err := r.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
placeholders := make([]string, 0, len(names)) placeholders := make([]string, 0, len(names))
args := make([]any, 0, len(names)) args := make([]any, 0, len(names))
@ -42,11 +36,12 @@ func (r *tagRepository) Upsert(names []string) error {
"SELECT name FROM input " + "SELECT name FROM input " +
"ON CONFLICT(name) DO NOTHING;" "ON CONFLICT(name) DO NOTHING;"
if _, err := tx.Exec(query, args...); err != nil { _, err := r.dbContext.ExecContext(ctx, query, args)
if err != nil {
return err return err
} }
return tx.Commit() return nil
} }
func (r *tagRepository) UpsertNoteTags(noteID int64, tagIDs []int64) error { func (r *tagRepository) UpsertNoteTags(noteID int64, tagIDs []int64) error {
@ -54,12 +49,6 @@ func (r *tagRepository) UpsertNoteTags(noteID int64, tagIDs []int64) error {
return nil return nil
} }
tx, err := r.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
placeholders := make([]string, 0, len(tagIDs)) placeholders := make([]string, 0, len(tagIDs))
args := make([]any, 0, len(tagIDs)) args := make([]any, 0, len(tagIDs))
@ -75,9 +64,10 @@ func (r *tagRepository) UpsertNoteTags(noteID int64, tagIDs []int64) error {
"ON CONFLICT(note_id, tag_id) DO NOTHING" + "ON CONFLICT(note_id, tag_id) DO NOTHING" +
"SELECT note_id, tag_id FROM input;" "SELECT note_id, tag_id FROM input;"
if _, err := tx.Exec(query, args...); err != nil { _, err := r.dbContext.ExecContext(context.Background(), query, args)
if err != nil {
return err return err
} }
return tx.Commit() return nil
} }

View file

@ -5,6 +5,12 @@ import (
"database/sql" "database/sql"
) )
type DBContext interface {
ExecContext(ctx context.Context, args ...any) (sql.Result, error)
QueryContext(ctx context.Context, args ...any) (*sql.Rows, error)
QueryRowContext(ctx context.Context, args ...any) *sql.Row
}
type UnitOfWork struct { type UnitOfWork struct {
db *sql.DB db *sql.DB
} }
@ -13,22 +19,22 @@ func NewUnitOfWork(db *sql.DB) *UnitOfWork {
return &UnitOfWork{db: db} return &UnitOfWork{db: db}
} }
func (u *UnitOfWork) Execute(ctx context.Context, fn func(tx *sql.Tx) error) error { func (u *UnitOfWork) Execute(ctx context.Context, fn func(tx *sql.Tx) error) (err error) {
tx, err := u.db.BeginTx(ctx, nil) transcation, err := u.db.BeginTx(ctx, nil)
if err != nil { if err != nil {
return err return err
} }
defer func() { defer func() {
if err != nil { if err != nil {
_ = tx.Rollback() _ = transcation.Rollback()
} }
}() }()
if err := fn(tx); err != nil { if err = fn(transcation); err != nil {
tx.Rollback()
return err return err
} }
return tx.Commit() err = transcation.Commit()
return err
} }