feat(hooks): add OnTxBegin and runInTx for resolvespec and restheadspec

Every transaction the handlers open now fires OnTxBegin first, with the
transaction in hookCtx.Tx, via common.RunRequestTx.
This commit is contained in:
2026-09-30 22:40:06 +02:00
parent cd96404cdd
commit f2dbe2561c
8 changed files with 313 additions and 51 deletions
+27
View File
@@ -0,0 +1,27 @@
package common
import "context"
// TxHookName is the shared value of every spec's OnTxBegin HookType.
const TxHookName = "on_tx_begin"
// TxContext is implemented by a spec's HookContext so RunRequestTx can point
// it at the transaction it opens.
type TxContext interface {
SetTx(tx Database)
}
// RunRequestTx opens a transaction on db, points tc at it, runs onBegin (the
// spec's OnTxBegin hooks) and then body. An error from onBegin or body rolls
// the transaction back; body is not run when onBegin fails.
func RunRequestTx(ctx context.Context, db Database, tc TxContext, onBegin func() error, body func(tx Database) error) error {
return db.RunInTransaction(ctx, func(tx Database) error {
tc.SetTx(tx)
if onBegin != nil {
if err := onBegin(); err != nil {
return err
}
}
return body(tx)
})
}
+36 -14
View File
@@ -313,7 +313,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
errMsg string
)
txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error {
txErr := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "read", options, w), func(tx common.Database) error {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
@@ -734,7 +734,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
if h.shouldUseNestedProcessor(v, model) {
logger.Info("Using nested CUD processor for create operation")
var nestedResult *common.ProcessResult
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
@@ -782,7 +782,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
// Standard processing without nested relations
pkName := reflection.GetPrimaryKeyName(model)
var responseData interface{} = v
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error {
hookCtx := &HookContext{
Context: ctx,
Handler: h,
@@ -857,7 +857,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
if hasNestedData {
logger.Info("Using nested CUD processor for batch create with nested data")
results := make([]map[string]interface{}, 0, len(v))
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error {
// Temporarily swap the database to use transaction
originalDB := h.nestedProcessor
h.nestedProcessor = common.NewNestedCUDProcessor(tx, h.registry, h)
@@ -912,7 +912,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
pkName := reflection.GetPrimaryKeyName(model)
modelElemType := reflection.GetPointerElement(reflect.TypeOf(model))
responseItems := make([]interface{}, 0, len(v))
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error {
for _, item := range v {
hookCtx := &HookContext{
Context: ctx,
@@ -989,7 +989,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
if hasNestedData {
logger.Info("Using nested CUD processor for batch create with nested data ([]interface{})")
results := make([]interface{}, 0, len(v))
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error {
// Temporarily swap the database to use transaction
originalDB := h.nestedProcessor
h.nestedProcessor = common.NewNestedCUDProcessor(tx, h.registry, h)
@@ -1046,7 +1046,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
pkName := reflection.GetPrimaryKeyName(model)
modelElemType := reflection.GetPointerElement(reflect.TypeOf(model))
responseItems := make([]interface{}, 0, len(v))
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "create", options, w), func(tx common.Database) error {
for _, item := range v {
itemMap, ok := item.(map[string]interface{})
if !ok {
@@ -1180,7 +1180,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
}
// Wrap in transaction to ensure BeforeUpdate hook is inside transaction
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
// Execute BeforeUpdate hooks inside transaction, before any queries run.
// BeforeUpdate hooks may set session-scoped RLS GUCs (via SET LOCAL);
// they must run before the existence-check select so that select is
@@ -1334,7 +1334,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
if hasNestedData {
logger.Info("Using nested CUD processor for batch update with nested data")
results := make([]map[string]interface{}, 0, len(updates))
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
// Temporarily swap the database to use transaction
originalDB := h.nestedProcessor
h.nestedProcessor = common.NewNestedCUDProcessor(tx, h.registry, h)
@@ -1368,7 +1368,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// Standard batch update without nested relations
pkName := reflection.GetPrimaryKeyName(model)
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
for _, item := range updates {
if itemID, ok := item["id"]; ok {
itemIDStr := fmt.Sprintf("%v", itemID)
@@ -1479,7 +1479,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
if hasNestedData {
logger.Info("Using nested CUD processor for batch update with nested data ([]interface{})")
results := make([]interface{}, 0, len(updates))
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
// Temporarily swap the database to use transaction
originalDB := h.nestedProcessor
h.nestedProcessor = common.NewNestedCUDProcessor(tx, h.registry, h)
@@ -1516,7 +1516,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// Standard batch update without nested relations
pkName := reflection.GetPrimaryKeyName(model)
list := make([]interface{}, 0)
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
for _, item := range updates {
if itemMap, ok := item.(map[string]interface{}); ok {
if itemID, ok := itemMap["id"]; ok {
@@ -1657,8 +1657,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
// state set by hooks (e.g. RLS settings) applies to every statement.
var payload interface{}
var failure *deleteFailure
txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error {
hookCtx.Tx = tx
txErr := h.runInTx(ctx, hookCtx, func(tx common.Database) error {
payload, failure = h.executeDelete(ctx, tx, hookCtx, schema, tableName, model, id, data)
if failure != nil {
return failure
@@ -2561,3 +2560,26 @@ func mergeWithInput(dbRecord interface{}, input map[string]interface{}) map[stri
}
return result
}
// newTxHookContext builds the context OnTxBegin hooks receive for paths that
// create their per-item hook contexts inside the transaction.
func (h *Handler) newTxHookContext(ctx context.Context, schema, entity string, model interface{}, operation string, options common.RequestOptions, w common.ResponseWriter) *HookContext {
return &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: operation,
Options: options,
Writer: w,
}
}
// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin
// fired first. Every transaction the handler opens goes through here.
func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(tx common.Database) error) error {
return common.RunRequestTx(ctx, h.db, hookCtx, func() error {
return h.hooks.Execute(OnTxBegin, hookCtx)
}, body)
}
+9
View File
@@ -41,6 +41,12 @@ const (
// Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each
// individual SQL-operation hook point, so it runs once per statement executed.
BeforeOp HookType = "before_op"
// OnTxBegin fires once, first, inside every transaction the handler opens
// (including the second short transaction for post-commit work). hookCtx.Tx
// is the transaction; use it to stamp transaction-local state such as RLS
// settings. An error or abort rolls the transaction back.
OnTxBegin HookType = common.TxHookName
)
// HookContext contains all the data available to a hook
@@ -76,6 +82,9 @@ type HookContext struct {
Tx common.Database
}
// SetTx points the context at the transaction in use (common.TxContext).
func (c *HookContext) SetTx(tx common.Database) { c.Tx = tx }
// HookFunc is the signature for hook functions
// It receives a HookContext and can modify it or return an error
// If an error is returned, the operation will be aborted
+84
View File
@@ -0,0 +1,84 @@
package resolvespec
import (
"errors"
"net/http"
"testing"
"github.com/DATA-DOG/go-sqlmock"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
// recordTxOrder records the order hooks fire in and the Tx each one saw.
func recordTxOrder(h *Handler, beginErr error) (*[]string, *[]common.Database) {
var order []string
var txs []common.Database
h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error {
order = append(order, "begin")
txs = append(txs, ctx.Tx)
return beginErr
})
h.Hooks().Register(BeforeDelete, func(ctx *HookContext) error {
order = append(order, "before_delete")
return nil
})
return &order, &txs
}
func TestOnTxBeginFiresOnceFirstOnTx(t *testing.T) {
cases := map[string]struct {
id string
data interface{}
exec int
}{
"single": {id: "7", exec: 1},
"batch": {data: []interface{}{"1", "2"}, exec: 2},
}
for name, tc := range cases {
t.Run(name, func(t *testing.T) {
h, mock, _ := newDeleteHarness(t)
order, txs := recordTxOrder(h, nil)
mock.ExpectBegin()
if tc.id != "" {
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
}
for i := 0; i < tc.exec; i++ {
mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1))
}
mock.ExpectCommit()
if rec := runDelete(h, tc.id, tc.data); rec.Code != http.StatusOK {
t.Fatalf("status %d body %s", rec.Code, rec.Body)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
if len(*txs) != 1 || (*txs)[0] == nil || (*txs)[0] == h.db {
t.Fatalf("OnTxBegin must fire once on the transaction, got %v", *txs)
}
if (*order)[0] != "begin" {
t.Fatalf("OnTxBegin must fire before other hooks, got %v", *order)
}
})
}
}
func TestOnTxBeginErrorRollsBackWithoutQueries(t *testing.T) {
h, mock, _ := newDeleteHarness(t)
order, _ := recordTxOrder(h, errors.New("no user"))
mock.ExpectBegin()
mock.ExpectRollback()
if rec := runDelete(h, "7", nil); rec.Code < http.StatusBadRequest {
t.Fatalf("status %d body %s", rec.Code, rec.Body)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
if len(*order) != 1 {
t.Fatalf("no hook may run after a failed OnTxBegin, got %v", *order)
}
}
+44 -26
View File
@@ -460,8 +460,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
errMsg string
)
txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error {
hookCtx.Tx = tx
txErr := h.runInTx(ctx, hookCtx, func(tx common.Database) error {
if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil {
statusCode, errCode, errMsg = http.StatusBadRequest, "hook_error", "Hook execution failed"
@@ -1322,8 +1321,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
// Process all items in a transaction
results := make([]interface{}, 0)
txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error {
hookCtx.Tx = tx
txErr := h.runInTx(ctx, hookCtx, func(tx common.Database) error {
if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil {
statusCode, errCode, errMsg = http.StatusBadRequest, "hook_error", "Hook execution failed"
return err
@@ -1538,11 +1536,23 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
// Variable to store the updated record
var updatedRecord interface{}
// Declare hook context to be used inside and outside transaction
var hookCtx *HookContext
// Hook context used inside and outside transaction
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
TableName: tableName,
Model: model,
Operation: "update",
Options: options,
ID: id,
Data: dataMap,
Writer: w,
}
// Process nested relations if present
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, hookCtx, func(tx common.Database) error {
// Create temporary nested processor with transaction
txNestedProcessor := common.NewNestedCUDProcessor(tx, h.registry, h)
@@ -1550,21 +1560,6 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
// BeforeUpdate hooks may set session-scoped RLS GUCs (via SET LOCAL);
// they must run before the existence-check select so that select is
// also subject to RLS on this connection/transaction.
hookCtx = &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
TableName: tableName,
Tx: tx,
Model: model,
Operation: "update",
Options: options,
ID: id,
Data: dataMap,
Writer: w,
}
if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil {
return fmt.Errorf("BeforeUpdate hook failed: %w", err)
}
@@ -1733,7 +1728,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
// Array of IDs as strings
logger.Info("Batch delete with %d IDs ([]string)", len(v))
deletedCount := 0
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, tableName, model, "delete", w), func(tx common.Database) error {
for _, itemID := range v {
// Execute hooks for each item
hookCtx := &HookContext{
@@ -1790,7 +1785,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
logger.Info("Batch delete with %d items ([]interface{})", len(v))
deletedCount := 0
pkName := reflection.GetPrimaryKeyName(model)
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, tableName, model, "delete", w), func(tx common.Database) error {
for _, item := range v {
var itemID interface{}
@@ -1864,7 +1859,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
logger.Info("Batch delete with %d items ([]map[string]interface{})", len(v))
deletedCount := 0
pkName := reflection.GetPrimaryKeyName(model)
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, tableName, model, "delete", w), func(tx common.Database) error {
for _, item := range v {
if itemID, ok := item[pkName]; ok && itemID != nil {
itemIDStr := fmt.Sprintf("%v", itemID)
@@ -1943,7 +1938,7 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id
// Lookup, hooks and delete share one transaction so transaction-local
// state set by hooks (e.g. RLS settings) applies to every statement.
var failure *deleteFailure
txErr := h.db.RunInTransaction(ctx, func(tx common.Database) error {
txErr := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, tableName, model, "delete", w), func(tx common.Database) error {
failure = h.deleteSingleInTx(ctx, tx, w, schema, entity, tableName, model, pkName, id, recordToDelete)
if failure != nil {
return failure
@@ -3532,3 +3527,26 @@ func (h *Handler) HandleOpenAPI(w common.ResponseWriter, r common.Request) {
func (h *Handler) SetOpenAPIGenerator(generator func() (string, error)) {
h.openAPIGenerator = generator
}
// newTxHookContext builds the context OnTxBegin hooks receive for paths that
// create their per-item hook contexts inside the transaction.
func (h *Handler) newTxHookContext(ctx context.Context, schema, entity, tableName string, model interface{}, operation string, w common.ResponseWriter) *HookContext {
return &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
TableName: tableName,
Model: model,
Operation: operation,
Writer: w,
}
}
// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin
// fired first. Every transaction the handler opens goes through here.
func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(tx common.Database) error) error {
return common.RunRequestTx(ctx, h.db, hookCtx, func() error {
return h.hooks.Execute(OnTxBegin, hookCtx)
}, body)
}
+9
View File
@@ -41,6 +41,12 @@ const (
// Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp fires at each
// individual SQL-operation hook point, so it runs once per statement executed.
BeforeOp HookType = "before_op"
// OnTxBegin fires once, first, inside every transaction the handler opens
// (including the second short transaction for post-commit work). hookCtx.Tx
// is the transaction; use it to stamp transaction-local state such as RLS
// settings. An error or abort rolls the transaction back.
OnTxBegin HookType = common.TxHookName
)
// HookContext contains all the data available to a hook
@@ -83,6 +89,9 @@ type HookContext struct {
Tx common.Database
}
// SetTx points the context at the transaction in use (common.TxContext).
func (c *HookContext) SetTx(tx common.Database) { c.Tx = tx }
// HookFunc is the signature for hook functions
// It receives a HookContext and can modify it or return an error
// If an error is returned, the operation will be aborted
+84
View File
@@ -0,0 +1,84 @@
package restheadspec
import (
"errors"
"net/http"
"testing"
"github.com/DATA-DOG/go-sqlmock"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
// recordTxOrder records the order hooks fire in and the Tx each one saw.
func recordTxOrder(h *Handler, beginErr error) (*[]string, *[]common.Database) {
var order []string
var txs []common.Database
h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error {
order = append(order, "begin")
txs = append(txs, ctx.Tx)
return beginErr
})
h.Hooks().Register(BeforeDelete, func(ctx *HookContext) error {
order = append(order, "before_delete")
return nil
})
return &order, &txs
}
func TestOnTxBeginFiresOnceFirstOnTx(t *testing.T) {
cases := map[string]struct {
id string
data interface{}
exec int
}{
"single": {id: "7", exec: 1},
"batch": {data: []interface{}{"1", "2"}, exec: 2},
}
for name, tc := range cases {
t.Run(name, func(t *testing.T) {
h, mock, _ := newDeleteHarness(t)
order, txs := recordTxOrder(h, nil)
mock.ExpectBegin()
if tc.id != "" {
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
}
for i := 0; i < tc.exec; i++ {
mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1))
}
mock.ExpectCommit()
if rec := runDelete(h, tc.id, tc.data); rec.Code != http.StatusOK {
t.Fatalf("status %d body %s", rec.Code, rec.Body)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
if len(*txs) != 1 || (*txs)[0] == nil || (*txs)[0] == h.db {
t.Fatalf("OnTxBegin must fire once on the transaction, got %v", *txs)
}
if (*order)[0] != "begin" {
t.Fatalf("OnTxBegin must fire before other hooks, got %v", *order)
}
})
}
}
func TestOnTxBeginErrorRollsBackWithoutQueries(t *testing.T) {
h, mock, _ := newDeleteHarness(t)
order, _ := recordTxOrder(h, errors.New("no user"))
mock.ExpectBegin()
mock.ExpectRollback()
if rec := runDelete(h, "7", nil); rec.Code < http.StatusBadRequest {
t.Fatalf("status %d body %s", rec.Code, rec.Body)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
if len(*order) != 1 {
t.Fatalf("no hook may run after a failed OnTxBegin, got %v", *order)
}
}