feat(tx): run update re-fetch and post-commit hooks in a second transaction

Re-fetch, BeforeScan and AfterUpdate/AfterCreate now run on a short
transaction that fires OnTxBegin. Existence selects inside the first
transaction use tx instead of the pool.
This commit is contained in:
2026-09-30 22:42:33 +02:00
parent eb492d52aa
commit ce706bacda
5 changed files with 278 additions and 79 deletions
+52 -37
View File
@@ -1210,7 +1210,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// Now read the existing record from the database
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...)
selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...)
// Apply conditions to select, based on the resolved target ID
// (URL ID, request ID, or the "id" field embedded in the data payload).
@@ -1292,22 +1292,25 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
return
}
// Fetch the updated record after the transaction commits to capture any trigger changes
// Fetch the updated record in a second short transaction after the first
// commit to capture any trigger changes (OnTxBegin re-applies RLS state).
updatedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
fetchQuery := h.db.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...)
if urlID != "" {
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID)
} else if reqID != nil {
switch id := reqID.(type) {
case string:
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
case []string:
if len(id) > 0 {
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s IN (?)", common.QuoteIdent(pkName)), id)
if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
fetchQuery := tx.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...)
if urlID != "" {
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID)
} else if reqID != nil {
switch id := reqID.(type) {
case string:
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
case []string:
if len(id) > 0 {
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s IN (?)", common.QuoteIdent(pkName)), id)
}
}
}
}
if err := fetchQuery.ScanModel(ctx); err != nil {
return fetchQuery.ScanModel(ctx)
}); err != nil {
logger.Error("Failed to fetch updated record: %v", err)
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
return
@@ -1375,7 +1378,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// First, read the existing record
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if err := selectQuery.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows {
continue // Skip if record not found
@@ -1441,19 +1444,25 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
return
}
// Fetch updated records after the transaction commits to capture any trigger changes
// Fetch updated records in a second short transaction after the first commit
// to capture any trigger changes (OnTxBegin re-applies RLS state).
fetchedUpdates := make([]interface{}, 0, len(updates))
for _, item := range updates {
if itemID, ok := item["id"]; ok && itemID != nil {
fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
fetchQuery := h.db.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if err := fetchQuery.ScanModel(ctx); err != nil {
logger.Error("Failed to fetch updated record with ID %v: %v", itemID, err)
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
return
if 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 && itemID != nil {
fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
fetchQuery := tx.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if err := fetchQuery.ScanModel(ctx); err != nil {
return fmt.Errorf("fetch updated record with ID %v: %w", itemID, err)
}
fetchedUpdates = append(fetchedUpdates, fetchedRecord)
}
fetchedUpdates = append(fetchedUpdates, fetchedRecord)
}
return nil
}); err != nil {
logger.Error("Failed to fetch updated records: %v", err)
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
return
}
logger.Info("Successfully updated %d records", len(fetchedUpdates))
@@ -1524,7 +1533,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// First, read the existing record
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
selectQuery := h.db.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
selectQuery := tx.NewSelect().Model(existingRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if err := selectQuery.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows {
continue // Skip if record not found
@@ -1593,21 +1602,27 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
return
}
// Fetch updated records after the transaction commits to capture any trigger changes
// Fetch updated records in a second short transaction after the first commit
// to capture any trigger changes (OnTxBegin re-applies RLS state).
fetchedList := make([]interface{}, 0, len(list))
for _, item := range list {
if itemMap, ok := item.(map[string]interface{}); ok {
if itemID, ok := itemMap["id"]; ok && itemID != nil {
fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
fetchQuery := h.db.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if err := fetchQuery.ScanModel(ctx); err != nil {
logger.Error("Failed to fetch updated record with ID %v: %v", itemID, err)
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
return
if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
for _, item := range list {
if itemMap, ok := item.(map[string]interface{}); ok {
if itemID, ok := itemMap["id"]; ok && itemID != nil {
fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
fetchQuery := tx.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if err := fetchQuery.ScanModel(ctx); err != nil {
return fmt.Errorf("fetch updated record with ID %v: %w", itemID, err)
}
fetchedList = append(fetchedList, fetchedRecord)
}
fetchedList = append(fetchedList, fetchedRecord)
}
}
return nil
}); err != nil {
logger.Error("Failed to fetch updated records: %v", err)
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
return
}
logger.Info("Successfully updated %d records", len(fetchedList))
+48
View File
@@ -0,0 +1,48 @@
package resolvespec
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) {
h, mock, _ := newDeleteHarness(t)
var begins []common.Database
h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error {
begins = append(begins, ctx.Tx)
return nil
})
cols := []string{"id", "name"}
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a"))
mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b"))
mock.ExpectCommit()
rec := httptest.NewRecorder()
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPut, "/", nil))
base, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
ctx := WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{})
h.handleUpdate(ctx, w, "7", nil, map[string]interface{}{"name": "b"}, common.RequestOptions{})
if 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(begins) != 2 || begins[0] == begins[1] {
t.Fatalf("OnTxBegin must fire once per transaction (2), got %d", len(begins))
}
}
+50 -39
View File
@@ -1459,10 +1459,8 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
}
}
// Execute AfterCreate hooks (runs after the transaction commits, against the
// pooled db — hookCtx.Tx was pointed at the now-closed transaction inside the
// RunInTransaction closure above and must not be reused here).
hookCtx.Tx = h.db
// Execute AfterCreate hooks in a second short transaction (the first has
// committed); OnTxBegin re-applies transaction-local state to it.
var responseData interface{}
if len(mergedResults) == 1 {
responseData = mergedResults[0]
@@ -1473,7 +1471,9 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
}
hookCtx.Error = nil
if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil {
if err := h.runInTx(ctx, hookCtx, func(common.Database) error {
return h.hooks.Execute(AfterCreate, hookCtx)
}); err != nil {
logger.Error("AfterCreate hook failed: %v", err)
h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
return
@@ -1571,7 +1571,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
// Now read the existing record from the database
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
selectQuery := h.db.NewSelect().Model(existingRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
selectQuery := tx.NewSelect().Model(existingRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
if err := selectQuery.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows {
return fmt.Errorf("record not found with ID: %v", targetID)
@@ -1657,43 +1657,54 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
return
}
// Fetch the updated record after the transaction commits to capture any trigger changes
fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface()
selectQuery := h.db.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
// Second short transaction: fetch the updated record after the first commit to
// capture any trigger changes, then run AfterUpdate. OnTxBegin re-applies
// transaction-local state (e.g. RLS settings) to this transaction.
var mergedData interface{}
var errCode, errMsg string
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface()
selectQuery := tx.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
// Execute BeforeScan hooks so row security is re-applied to the post-update
// re-fetch, same as it is for the initial read and the update query itself.
// Without this, the re-fetch can return a row the caller isn't authorized to see.
// The transaction has already committed by this point, so hooks must use the
// pooled connection rather than the now-dead tx.
hookCtx.Tx = h.db
hookCtx.Query = selectQuery
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
logger.Error("BeforeScan hook failed: %v", err)
h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
return
}
if modifiedQuery, ok := hookCtx.Query.(common.SelectQuery); ok {
selectQuery = modifiedQuery
}
// Execute BeforeScan hooks so row security is re-applied to the post-update
// re-fetch, same as it is for the initial read and the update query itself.
// Without this, the re-fetch can return a row the caller isn't authorized to see.
hookCtx.Query = selectQuery
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
logger.Error("BeforeScan hook failed: %v", err)
errCode, errMsg = "hook_error", "Hook execution failed"
return err
}
if modifiedQuery, ok := hookCtx.Query.(common.SelectQuery); ok {
selectQuery = modifiedQuery
}
if err := selectQuery.ScanModel(ctx); err != nil {
logger.Error("Failed to fetch updated record: %v", err)
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
return
}
updatedRecord = fetchedRecord
if err := selectQuery.ScanModel(ctx); err != nil {
logger.Error("Failed to fetch updated record: %v", err)
errCode, errMsg = "fetch_error", "Failed to fetch updated record"
return err
}
updatedRecord = fetchedRecord
// Merge the updated record with the original request data
// This preserves extra keys from the request and updates values from the database
mergedData := h.mergeRecordWithRequest(updatedRecord, dataMap)
// Merge the updated record with the original request data
// This preserves extra keys from the request and updates values from the database
mergedData = h.mergeRecordWithRequest(updatedRecord, dataMap)
// Execute AfterUpdate hooks
hookCtx.Result = mergedData
hookCtx.Error = nil
if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil {
logger.Error("AfterUpdate hook failed: %v", err)
h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
// Execute AfterUpdate hooks
hookCtx.Result = mergedData
hookCtx.Error = nil
if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil {
logger.Error("AfterUpdate hook failed: %v", err)
errCode, errMsg = "hook_error", "Hook execution failed"
return err
}
return nil
})
if err != nil {
if errCode == "" {
errCode, errMsg = "fetch_error", "Failed to fetch updated record"
}
h.sendError(w, http.StatusInternalServerError, errCode, errMsg, err)
return
}
+123
View File
@@ -0,0 +1,123 @@
package restheadspec
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/pgdialect"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
func TestUpdateRefetchAndAfterUpdateRunInSecondTransaction(t *testing.T) {
// The bun adapter builds model-based updates; the pgsql adapter does not.
sqlDB, mock, err := sqlmock.New()
if err != nil {
t.Fatal(err)
}
sqlDB.SetMaxOpenConns(1)
t.Cleanup(func() { _ = sqlDB.Close() })
h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry())
var begins []common.Database
var order []string
var afterTx common.Database
h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error {
begins = append(begins, ctx.Tx)
order = append(order, "begin")
return nil
})
h.Hooks().Register(AfterUpdate, func(ctx *HookContext) error {
afterTx = ctx.Tx
order = append(order, "after_update")
return nil
})
cols := []string{"id", "name"}
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a"))
mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b"))
mock.ExpectCommit()
rec := httptest.NewRecorder()
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPut, "/", nil))
base, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
ctx := WithSchema(base, "public")
ctx = WithEntity(ctx, "items")
ctx = WithTableName(ctx, "items")
ctx = WithModel(ctx, delItem{})
h.handleUpdate(ctx, w, "7", nil, map[string]interface{}{"name": "b"}, ExtendedRequestOptions{})
if 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(begins) != 2 || begins[0] == begins[1] {
t.Fatalf("OnTxBegin must fire once per transaction (2), got %d", len(begins))
}
if afterTx == nil || afterTx == h.db || afterTx != begins[1] {
t.Fatalf("AfterUpdate must run on the second transaction")
}
if len(order) != 3 || order[2] != "after_update" {
t.Fatalf("unexpected hook order %v", order)
}
}
func TestAfterCreateRunsInSecondTransaction(t *testing.T) {
sqlDB, mock, err := sqlmock.New()
if err != nil {
t.Fatal(err)
}
sqlDB.SetMaxOpenConns(1)
t.Cleanup(func() { _ = sqlDB.Close() })
h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry())
var begins []common.Database
var afterTx common.Database
h.Hooks().Register(OnTxBegin, func(ctx *HookContext) error {
begins = append(begins, ctx.Tx)
return nil
})
h.Hooks().Register(AfterCreate, func(ctx *HookContext) error {
afterTx = ctx.Tx
return nil
})
mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectCommit()
rec := httptest.NewRecorder()
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil))
base, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
ctx := WithSchema(base, "public")
ctx = WithEntity(ctx, "items")
ctx = WithTableName(ctx, "items")
ctx = WithModel(ctx, delItem{})
h.handleCreate(ctx, w, map[string]interface{}{"name": "a"}, ExtendedRequestOptions{})
if 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(begins) != 2 || afterTx == nil || afterTx == h.db || afterTx != begins[1] {
t.Fatalf("AfterCreate must run on the second transaction, begins=%d", len(begins))
}
}