From f070674cdbde30da7744c12795dce439e2874803 Mon Sep 17 00:00:00 2001 From: Kristian Borgwarth <10348902@pm.me> Date: Thu, 2 Apr 2026 17:24:18 +0200 Subject: [PATCH] feat(uow): transaction share by retun tx --- core/handlers/create_note_handler.go | 8 +--- persistence/repositories/note_repository.go | 9 +++-- persistence/repositories/tag_repository.go | 11 ++--- persistence/repositories/uow.go | 45 ++++++++++----------- 4 files changed, 35 insertions(+), 38 deletions(-) diff --git a/core/handlers/create_note_handler.go b/core/handlers/create_note_handler.go index 1fe121d..81a443d 100644 --- a/core/handlers/create_note_handler.go +++ b/core/handlers/create_note_handler.go @@ -17,13 +17,9 @@ type createNoteCommand struct { Vars map[string]string `json:"vars"` } -type CreateNoteHandler struct { - uow *repositories.UnitOfWork -} +type CreateNoteHandler struct {} -func NewCreateNoteHandler(uow *repositories.UnitOfWork) *CreateNoteHandler { - return &CreateNoteHandler{uow: uow} -} +func NewCreateNoteHandler(uow *repositories.UnitOfWork) *CreateNoteHandler func (h CreateNoteHandler) Handle(ctx context.Context, raw json.RawMessage) (any, error) { var cmd createNoteCommand diff --git a/persistence/repositories/note_repository.go b/persistence/repositories/note_repository.go index 0c858b7..3ad8b06 100644 --- a/persistence/repositories/note_repository.go +++ b/persistence/repositories/note_repository.go @@ -2,6 +2,7 @@ package repositories import ( "context" + "database/sql" ) type NoteRepository interface { @@ -9,11 +10,11 @@ type NoteRepository interface { } type noteRepository struct { - dbCtx DBContext + Transaction *sql.Tx } -func NewNoteRepository(dbContext DBContext) NoteRepository { - return ¬eRepository{dbCtx: dbContext} +func NewNoteRepository(tx *sql.Tx) NoteRepository { + return ¬eRepository{Transaction: tx} } func (r *noteRepository) Upsert(ctx context.Context, title string, path string, slug string) error { @@ -24,6 +25,6 @@ func (r *noteRepository) Upsert(ctx context.Context, title string, path string, SET title = EXCLUDED.title, path = EXCLUDED.path; ` - _, err := r.dbCtx.ExecContext(ctx, query, title, path, slug) + _, err := r.Transaction.ExecContext(ctx, query, title, path, slug) return err } diff --git a/persistence/repositories/tag_repository.go b/persistence/repositories/tag_repository.go index 3f8b950..f1446cf 100644 --- a/persistence/repositories/tag_repository.go +++ b/persistence/repositories/tag_repository.go @@ -2,6 +2,7 @@ package repositories import ( "context" + "database/sql" "strings" ) @@ -10,11 +11,11 @@ type TagRepository interface { } type tagRepository struct { - dbContext DBContext + Transaction *sql.Tx } -func NewTagRepository(dbContext DBContext) TagRepository { - return &tagRepository{dbContext: dbContext} +func NewTagRepository(tx *sql.Tx) TagRepository { +return &tagRepository{Transaction: tx} } func (r *tagRepository) Upsert(ctx context.Context, names []string) error { @@ -36,7 +37,7 @@ func (r *tagRepository) Upsert(ctx context.Context, names []string) error { "SELECT name FROM input " + "ON CONFLICT(name) DO NOTHING;" - _, err := r.dbContext.ExecContext(ctx, query, args) + _, err := r.Transaction.ExecContext(ctx, query, args) if err != nil { return err } @@ -64,7 +65,7 @@ func (r *tagRepository) UpsertNoteTags(noteID int64, tagIDs []int64) error { "ON CONFLICT(note_id, tag_id) DO NOTHING" + "SELECT note_id, tag_id FROM input;" - _, err := r.dbContext.ExecContext(context.Background(), query, args) + _, err := r.Transaction.ExecContext(context.Background(), query, args) if err != nil { return err } diff --git a/persistence/repositories/uow.go b/persistence/repositories/uow.go index 4213763..cbec578 100644 --- a/persistence/repositories/uow.go +++ b/persistence/repositories/uow.go @@ -1,40 +1,39 @@ + package repositories import ( - "context" "database/sql" ) -type DBContext 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 -} - type UnitOfWork struct { db *sql.DB + Transaction *sql.Tx } func NewUnitOfWork(db *sql.DB) *UnitOfWork { return &UnitOfWork{db: db} } -func (u *UnitOfWork) Execute(ctx context.Context, fn func(tx *sql.Tx) error) (err error) { - transcation, err := u.db.BeginTx(ctx, nil) +func (u *UnitOfWork) Begin() (tx *sql.Tx, err error) { + tx, err = u.db.Begin() if err != nil { - return err + return nil, err } - - defer func() { - if err != nil { - _ = transcation.Rollback() - } - }() - - if err = fn(transcation); err != nil { - return err - } - - err = transcation.Commit() - return err + u.Transaction = tx + return tx, nil } + +func (u *UnitOfWork) Commit() error { + if u.Transaction == nil { + return nil + } + return u.Transaction.Commit() +} + +func (u *UnitOfWork) Rollback() error { + if u.Transaction == nil { + return nil + } + return u.Transaction.Rollback() +} +