| package entutils |
|
|
| import ( |
| "context" |
| "database/sql" |
| "fmt" |
| "strconv" |
| "sync" |
|
|
| "entgo.io/ent/dialect" |
|
|
| "github.com/openmeterio/openmeter/pkg/framework/transaction" |
| ) |
|
|
| type RawEntConfig struct { |
| |
| Driver dialect.Driver |
| |
| Debug bool |
| |
| Log func(...any) |
|
|
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| } |
|
|
| type Transactable interface { |
| Commit() error |
| Rollback() error |
| SavePoint(name string) error |
| RollbackTo(name string) error |
| Release(name string) error |
| } |
|
|
| type TxHijacker interface { |
| HijackTx(ctx context.Context, opts *sql.TxOptions) (context.Context, *RawEntConfig, Transactable, error) |
| } |
|
|
| func NewTxDriver(driver Transactable, cfg *RawEntConfig) *TxDriver { |
| return &TxDriver{ |
| driver: driver, |
| cfg: cfg, |
| } |
| } |
|
|
| type txSavepoint int |
|
|
| const ( |
| txSavepointNone txSavepoint = 0 |
| ) |
|
|
| func (sp txSavepoint) Next() txSavepoint { |
| return sp + 1 |
| } |
|
|
| func (sp txSavepoint) Prev() txSavepoint { |
| if sp == txSavepointNone { |
| return txSavepointNone |
| } |
|
|
| return sp - 1 |
| } |
|
|
| func (sp txSavepoint) String() string { |
| return "s" + strconv.Itoa(int(sp)) |
| } |
|
|
| type TxDriver struct { |
| driver Transactable |
| |
| |
| cfg *RawEntConfig |
|
|
| mu sync.Mutex |
| once sync.Once |
|
|
| currentSavepoint txSavepoint |
|
|
| err error |
| } |
|
|
| var _ transaction.Driver = &TxDriver{} |
|
|
| func (t *TxDriver) GetConfig() *RawEntConfig { |
| return t.cfg |
| } |
|
|
| |
| func (t *TxDriver) Commit() error { |
| |
| t.mu.Lock() |
| defer t.mu.Unlock() |
|
|
| |
| if t.err != nil { |
| return t.err |
| } |
|
|
| if t.currentSavepoint != txSavepointNone { |
| |
| if err := t.driver.Release(t.currentSavepoint.String()); err == nil { |
| t.currentSavepoint = t.currentSavepoint.Prev() |
| } else { |
| t.err = err |
| } |
| } else { |
| |
| t.err = t.driver.Commit() |
| } |
|
|
| return t.err |
| } |
|
|
| |
| func (t *TxDriver) Rollback() error { |
| |
| t.mu.Lock() |
| defer t.mu.Unlock() |
|
|
| |
| if t.err != nil { |
| return t.err |
| } |
|
|
| if t.currentSavepoint != txSavepointNone { |
| |
| if err := t.driver.RollbackTo(t.currentSavepoint.String()); err == nil { |
| t.currentSavepoint = t.currentSavepoint.Prev() |
| } else { |
| t.err = err |
| } |
| } else { |
| |
| t.err = t.driver.Rollback() |
| } |
|
|
| return t.err |
| } |
|
|
| func (t *TxDriver) SavePoint() error { |
| t.mu.Lock() |
| defer t.mu.Unlock() |
|
|
| skipSavePoint := false |
|
|
| t.once.Do(func() { |
| |
| |
| |
| skipSavePoint = true |
| }) |
|
|
| if !skipSavePoint { |
| next := t.currentSavepoint.Next() |
|
|
| err := t.driver.SavePoint(next.String()) |
| if err != nil { |
| return err |
| } |
|
|
| t.currentSavepoint = next |
| } |
|
|
| return nil |
| } |
|
|
| |
| type TxCreator = transaction.Creator |
|
|
| |
| type TxUser[T any] interface { |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| WithTx(ctx context.Context, tx *TxDriver) T |
| Self() T |
| } |
|
|
| |
| |
| func TransactingRepo[R, T any]( |
| ctx context.Context, |
| repo interface { |
| TxUser[T] |
| TxCreator |
| }, |
| cb func(ctx context.Context, rep T) (R, error), |
| ) (R, error) { |
| var def R |
| tx, err := GetDriverFromContext(ctx) |
| if err != nil { |
| |
| if _, ok := err.(*transaction.DriverNotFoundError); !ok { |
| return def, err |
| } |
|
|
| |
| return cb(ctx, repo.Self()) |
| } |
|
|
| |
| return cb(ctx, repo.WithTx(ctx, tx)) |
| } |
|
|
| |
| func TransactingRepoWithNoValue[T any]( |
| ctx context.Context, |
| repo interface { |
| TxUser[T] |
| TxCreator |
| }, |
| cb func(ctx context.Context, rep T) error, |
| ) error { |
| _, err := TransactingRepo(ctx, repo, func(ctx context.Context, rep T) (interface{}, error) { |
| return nil, cb(ctx, rep) |
| }) |
| return err |
| } |
|
|
| func asEntDriver(drv transaction.Driver) (*TxDriver, error) { |
| entTxDriver, ok := drv.(*TxDriver) |
| if !ok { |
| return nil, fmt.Errorf("tx driver is not ent tx driver") |
| } |
| return entTxDriver, nil |
| } |
|
|
| |
| func GetDriverFromContext(ctx context.Context) (*TxDriver, error) { |
| driver, err := transaction.GetDriverFromContext(ctx) |
| if err != nil { |
| return nil, err |
| } |
| return asEntDriver(driver) |
| } |
|
|