| package transaction |
|
|
| import ( |
| "context" |
| "errors" |
| "fmt" |
| "log/slog" |
| "runtime/debug" |
| ) |
|
|
| |
| type Driver interface { |
| Commit() error |
| Rollback() error |
| SavePoint() error |
| } |
|
|
| |
| type Creator interface { |
| Tx(ctx context.Context) (context.Context, Driver, error) |
| } |
|
|
| |
| func RunWithNoValue(ctx context.Context, creator Creator, cb func(ctx context.Context) error) error { |
| _, err := Run(ctx, creator, func(ctx context.Context) (interface{}, error) { |
| return nil, cb(ctx) |
| }) |
| return err |
| } |
|
|
| |
| func Run[R any](ctx context.Context, creator Creator, cb func(ctx context.Context) (R, error)) (R, error) { |
| var def R |
| |
| ctx, tx, err := getTx(ctx, creator) |
| if err != nil { |
| return def, err |
| } |
|
|
| |
| ctx, err = SetDriverOnContext(ctx, tx) |
| if _, ok := err.(*DriverConflictError); !ok && err != nil { |
| return def, fmt.Errorf("unknown error %w", err) |
| } |
|
|
| |
| return manage(ctx, tx, func(ctx context.Context, tx Driver) (R, error) { |
| return cb(ctx) |
| }) |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| func RunInNewTransaction[R any](ctx context.Context, creator Creator, cb func(ctx context.Context) (R, error)) (R, error) { |
| var def R |
|
|
| ctx, tx, err := creator.Tx(ctx) |
| if err != nil { |
| return def, fmt.Errorf("failed to start transaction: %w", err) |
| } |
|
|
| ctx = withDriver(ctx, tx) |
|
|
| return manage(ctx, tx, func(ctx context.Context, tx Driver) (R, error) { |
| return cb(ctx) |
| }) |
| } |
|
|
| |
| func getTx(ctx context.Context, creator Creator) (context.Context, Driver, error) { |
| if tx, err := GetDriverFromContext(ctx); err == nil { |
| return ctx, tx, nil |
| } else { |
| if _, ok := err.(*DriverNotFoundError); !ok { |
| slog.Debug("failed to get transaction from context", "transaction_error", err) |
| } |
| ctx, tx, err := creator.Tx(ctx) |
| if err != nil { |
| return nil, nil, fmt.Errorf("failed to start transaction: %w", err) |
| } |
| return ctx, tx, err |
| } |
| } |
|
|
| |
| func manage[R any](ctx context.Context, tx Driver, cb func(ctx context.Context, tx Driver) (R, error)) (R, error) { |
| var def R |
| defer func() { |
| if r := recover(); r != nil { |
| pMsg := fmt.Sprintf("%v:\n%s", r, debug.Stack()) |
|
|
| |
| _ = tx.Rollback() |
| panic(pMsg) |
| } |
| }() |
|
|
| err := tx.SavePoint() |
| if err != nil { |
| return def, err |
| } |
|
|
| result, err := cb(ctx, tx) |
| if err != nil { |
| |
| if rerr := tx.Rollback(); rerr != nil { |
| err = errors.Join(err, rerr) |
| } |
|
|
| return def, err |
| } |
|
|
| |
| err = tx.Commit() |
| if err != nil { |
| if rerr := tx.Rollback(); rerr != nil { |
| err = errors.Join(err, rerr) |
| } |
| return def, err |
| } |
|
|
| return result, nil |
| } |
|
|