mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 12:31:59 +00:00
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:
+52
-37
@@ -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))
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user