diff --git a/audit/single_tran.md b/audit/single_tran.md index 0463c33..aaecd87 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -69,20 +69,29 @@ - Consumer's ResolveSpec version: confirm it is >= v1.1.28 (read/create already in tx). Not blocking. ## Phases -| # | Change | Files | Notes | -|---|---|---|---| -| 0 | 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 | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS | -| 2 | `OnTxBegin` hook type + `runInTx` helper | `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec first | -| 3 | 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 | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias | -| 5 | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx | -| 6 | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | -| 7 | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | +| # | Status | Change | Files | Notes | +|---|---|---|---|---| +| 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 | +| 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 | +| 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 | +| 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | | +| 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | | + +## Progress +- DONE P0: baseline via `dbtrace` on real Postgres (commit `cd96404`): create/read/delete `pooled=0`; update `pooled=1` (re-fetch) = P3 target. websocketspec/mqttspec/resolvemcp not measured. +- DONE P1: single + batch delete in one tx (resolvespec, restheadspec). Not done: per-item `BeforeDelete` in resolvespec batch (behavior change, deferred). +- 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. +- 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`. +- NEXT: P3. ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. -- Missing: any test asserting hook `Tx` is a tx, or counting connections per op. +- Done: delete tx tests (`pkg/*/delete_tx_test.go`, sqlmock, 1-conn pool detects pool use). Missing: same for read/create/update, `OnTxBegin`, other specs. - Add per spec/op: hook `Tx` is not the pool; `OnTxBegin` fires once per tx, before other hooks; single-ID delete = 1 tx; `dbtrace` `pooled == 0` on the request path. - Test data: reuse `pkg/testmodels`; **ask before generating new data** (per project rule). - Regression: full `go test -race` for security, dbmanager, common, restheadspec, resolvespec, websocketspec, mqttspec, resolvemcp, funcspec. Known pre-existing failures: mqttspec integration (no DB). diff --git a/pkg/common/txhook.go b/pkg/common/txhook.go new file mode 100644 index 0000000..4d0023d --- /dev/null +++ b/pkg/common/txhook.go @@ -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) + }) +} diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index cdbabb2..9c7e9f7 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -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) +} diff --git a/pkg/resolvespec/hooks.go b/pkg/resolvespec/hooks.go index 5651e7a..184dd31 100644 --- a/pkg/resolvespec/hooks.go +++ b/pkg/resolvespec/hooks.go @@ -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 diff --git a/pkg/resolvespec/on_tx_begin_test.go b/pkg/resolvespec/on_tx_begin_test.go new file mode 100644 index 0000000..df0096f --- /dev/null +++ b/pkg/resolvespec/on_tx_begin_test.go @@ -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) + } +} diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index 39a8349..3daf210 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -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) +} diff --git a/pkg/restheadspec/hooks.go b/pkg/restheadspec/hooks.go index 57b2320..a0c628b 100644 --- a/pkg/restheadspec/hooks.go +++ b/pkg/restheadspec/hooks.go @@ -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 diff --git a/pkg/restheadspec/on_tx_begin_test.go b/pkg/restheadspec/on_tx_begin_test.go new file mode 100644 index 0000000..2b473e7 --- /dev/null +++ b/pkg/restheadspec/on_tx_begin_test.go @@ -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) + } +}