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
+5 -3
View File
@@ -74,7 +74,7 @@
| 0 | DONE | Baseline: enable `dbtrace` on testserver, record `tx/tx_queries/pooled/raw` per op | `cmd/testserver`, `pkg/dbtrace` | pooled > 0 on write ops = the gaps above | | 0 | DONE | Baseline: enable `dbtrace` on testserver, record `tx/tx_queries/pooled/raw` per op | `cmd/testserver`, `pkg/dbtrace` | pooled > 0 on write ops = the gaps above |
| 1 | DONE | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS | | 1 | DONE | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS |
| 2 | DONE | `OnTxBegin` hook type + `runInTx` helper | `common/txhook.go`, `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec; other specs in P4-6 | | 2 | DONE | `OnTxBegin` hook type + `runInTx` helper | `common/txhook.go`, `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec; other specs in P4-6 |
| 3 | TODO | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above | | 3 | DONE | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above |
| 4 | TODO | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | | 4 | TODO | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias |
| 5 | TODO | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | | 5 | TODO | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx |
| 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | | 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | |
@@ -86,8 +86,10 @@
- DONE infra: `sqlmock` delete tx tests (both specs); compose test server + `scripts/testserver-smoke.sh` (podman first); testmodels ids now serial. - DONE infra: `sqlmock` delete tx tests (both specs); compose test server + `scripts/testserver-smoke.sh` (podman first); testmodels ids now serial.
- NOTE: restheadspec single delete still does the lookup before `BeforeDelete`; safe once `OnTxBegin` (P2) exists. An `AfterDelete` failure now rolls the delete back. - NOTE: restheadspec single delete still does the lookup before `BeforeDelete`; safe once `OnTxBegin` (P2) exists. An `AfterDelete` failure now rolls the delete back.
- DONE P2 (resolvespec + restheadspec): `common.TxHookName`, `common.TxContext` (`SetTx` only; no abort/context accessors needed since `Execute` already returns an error on abort), `common.RunRequestTx`; per-spec `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Every `RunInTransaction` in both handlers now goes through it. Tests: `pkg/*/on_tx_begin_test.go` (once, first, on tx, failure rolls back). Not yet: the post-commit second tx (P3) and the security stamping registration (P7). - DONE P2 (resolvespec + restheadspec): `common.TxHookName`, `common.TxContext` (`SetTx` only; no abort/context accessors needed since `Execute` already returns an error on abort), `common.RunRequestTx`; per-spec `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Every `RunInTransaction` in both handlers now goes through it. Tests: `pkg/*/on_tx_begin_test.go` (once, first, on tx, failure rolls back). Not yet: the post-commit second tx (P3) and the security stamping registration (P7).
- FOUND (P3 scope): resolvespec batch update (`handler.go` ~`:1377`, `:1529`) reads existing record via `h.db.NewSelect()` inside the tx = pool connection; should be `tx`. - DONE P3: restheadspec update re-fetch + `BeforeScan` + `AfterUpdate` and `AfterCreate` run in a second short `runInTx`; resolvespec update re-fetches (single, both batch paths) run in a second short `runInTx`. Fixed the pool reads inside the first tx (resolvespec single/batch update existing-record select, restheadspec update existence select) to use `tx`. Tests: `pkg/*/update_tx_test.go` (restheadspec uses the bun adapter; the pgsql adapter cannot build model-based updates).
- NEXT: P3. - NOTE: resolvespec fires no `AfterCreate`/`AfterRead`/`AfterUpdate`-post-commit hooks other than `AfterUpdate` inside the tx; nothing more to move there.
- OPEN: restheadspec `AfterRead` still runs post-commit with `Tx = h.db` (`:1004`); decision says read has no second tx. Needs a call: run it inside the read tx, or in a short second tx.
- NEXT: P4.
## Tests ## Tests
- Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests.
+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 // Now read the existing record from the database
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() 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 // Apply conditions to select, based on the resolved target ID
// (URL ID, request ID, or the "id" field embedded in the data payload). // (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 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() updatedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
fetchQuery := h.db.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...) if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
if urlID != "" { fetchQuery := tx.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...)
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID) if urlID != "" {
} else if reqID != nil { fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID)
switch id := reqID.(type) { } else if reqID != nil {
case string: switch id := reqID.(type) {
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) case string:
case []string: fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
if len(id) > 0 { case []string:
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s IN (?)", common.QuoteIdent(pkName)), id) if len(id) > 0 {
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s IN (?)", common.QuoteIdent(pkName)), id)
}
} }
} }
} return fetchQuery.ScanModel(ctx)
if err := fetchQuery.ScanModel(ctx); err != nil { }); err != nil {
logger.Error("Failed to fetch updated record: %v", err) logger.Error("Failed to fetch updated record: %v", err)
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err)
return return
@@ -1375,7 +1378,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
// First, read the existing record // First, read the existing record
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() 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 := selectQuery.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
continue // Skip if record not found continue // Skip if record not found
@@ -1441,19 +1444,25 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
return 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)) fetchedUpdates := make([]interface{}, 0, len(updates))
for _, item := range updates { if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
if itemID, ok := item["id"]; ok && itemID != nil { for _, item := range updates {
fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() if itemID, ok := item["id"]; ok && itemID != nil {
fetchQuery := h.db.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
if err := fetchQuery.ScanModel(ctx); err != nil { fetchQuery := tx.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
logger.Error("Failed to fetch updated record with ID %v: %v", itemID, err) if err := fetchQuery.ScanModel(ctx); err != nil {
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) return fmt.Errorf("fetch updated record with ID %v: %w", itemID, err)
return }
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)) 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 // First, read the existing record
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() 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 := selectQuery.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
continue // Skip if record not found continue // Skip if record not found
@@ -1593,21 +1602,27 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
return 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)) fetchedList := make([]interface{}, 0, len(list))
for _, item := range list { if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
if itemMap, ok := item.(map[string]interface{}); ok { for _, item := range list {
if itemID, ok := itemMap["id"]; ok && itemID != nil { if itemMap, ok := item.(map[string]interface{}); ok {
fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() if itemID, ok := itemMap["id"]; ok && itemID != nil {
fetchQuery := h.db.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) fetchedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
if err := fetchQuery.ScanModel(ctx); err != nil { fetchQuery := tx.NewSelect().Model(fetchedRecord).Column(reflection.GetSQLModelColumns(model)...).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
logger.Error("Failed to fetch updated record with ID %v: %v", itemID, err) if err := fetchQuery.ScanModel(ctx); err != nil {
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) return fmt.Errorf("fetch updated record with ID %v: %w", itemID, err)
return }
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)) 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 // Execute AfterCreate hooks in a second short transaction (the first has
// pooled db — hookCtx.Tx was pointed at the now-closed transaction inside the // committed); OnTxBegin re-applies transaction-local state to it.
// RunInTransaction closure above and must not be reused here).
hookCtx.Tx = h.db
var responseData interface{} var responseData interface{}
if len(mergedResults) == 1 { if len(mergedResults) == 1 {
responseData = mergedResults[0] responseData = mergedResults[0]
@@ -1473,7 +1471,9 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
} }
hookCtx.Error = nil 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) logger.Error("AfterCreate hook failed: %v", err)
h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
return return
@@ -1571,7 +1571,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
// Now read the existing record from the database // Now read the existing record from the database
existingRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface() 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 := selectQuery.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
return fmt.Errorf("record not found with ID: %v", targetID) 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 return
} }
// Fetch the updated record after the transaction commits to capture any trigger changes // Second short transaction: fetch the updated record after the first commit to
fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface() // capture any trigger changes, then run AfterUpdate. OnTxBegin re-applies
selectQuery := h.db.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) // 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 // 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. // 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. // 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 hookCtx.Query = selectQuery
// pooled connection rather than the now-dead tx. if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil {
hookCtx.Tx = h.db logger.Error("BeforeScan hook failed: %v", err)
hookCtx.Query = selectQuery errCode, errMsg = "hook_error", "Hook execution failed"
if err := h.hooks.ExecuteBeforeOp(BeforeScan, hookCtx); err != nil { return err
logger.Error("BeforeScan hook failed: %v", err) }
h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err) if modifiedQuery, ok := hookCtx.Query.(common.SelectQuery); ok {
return selectQuery = modifiedQuery
} }
if modifiedQuery, ok := hookCtx.Query.(common.SelectQuery); ok {
selectQuery = modifiedQuery
}
if err := selectQuery.ScanModel(ctx); err != nil { if err := selectQuery.ScanModel(ctx); err != nil {
logger.Error("Failed to fetch updated record: %v", err) logger.Error("Failed to fetch updated record: %v", err)
h.sendError(w, http.StatusInternalServerError, "fetch_error", "Failed to fetch updated record", err) errCode, errMsg = "fetch_error", "Failed to fetch updated record"
return return err
} }
updatedRecord = fetchedRecord updatedRecord = fetchedRecord
// Merge the updated record with the original request data // Merge the updated record with the original request data
// This preserves extra keys from the request and updates values from the database // This preserves extra keys from the request and updates values from the database
mergedData := h.mergeRecordWithRequest(updatedRecord, dataMap) mergedData = h.mergeRecordWithRequest(updatedRecord, dataMap)
// Execute AfterUpdate hooks // Execute AfterUpdate hooks
hookCtx.Result = mergedData hookCtx.Result = mergedData
hookCtx.Error = nil hookCtx.Error = nil
if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil {
logger.Error("AfterUpdate hook failed: %v", err) logger.Error("AfterUpdate hook failed: %v", err)
h.sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", 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 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))
}
}