diff --git a/audit/single_tran.md b/audit/single_tran.md index d212f67..380f30e 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -96,7 +96,10 @@ - DONE (AfterRead, decided by user): restheadspec `AfterRead` now runs in a second short tx. Test: `pkg/restheadspec/read_tx_test.go`. - DONE P7: `pkg/security/txsettings.go`: `SecurityList.SetTxSettings(fn)`, `StampTxSettings`, `ApplyTxSettings` (configurable map, decided by user; `set_config(name, value, true)`, value hex-encoded, name validated, Postgres only, fail closed). Every spec's `RegisterSecurityHooks` registers it on `OnTxBegin`. Tests: `pkg/security/txsettings_test.go`, `pkg/resolvespec/tx_settings_test.go`. Docs: `pkg/common/TRANSACTIONS.md`. - DONE real-Postgres check (resolvespec, testserver via compose): create `tx=1 pooled=0`, read `tx=1 pooled=0`, update `tx=2 pooled=0` (was `pooled=1`), single delete `tx=1 pooled=0`, batch create/delete `tx=1 pooled=0`. Compose now uses host networking (bridge fails here): testserver on 8123, Postgres on 8124 (was 8080/5434); integration test DSNs updated. Smoke script covers read and update. websocketspec/mqttspec/resolvemcp/restheadspec/funcspec not measured on real Postgres. -- NEXT: extra per-spec create/update tests (optional). +- DONE regression tests: per-spec read/create/update/delete hook-on-tx, failure-rollback (Before*/After*/`OnTxBegin`) and second-tx tests in all six specs (`ops_tx_test.go`, `tx_test.go`, `read_tx_test.go`); stamping tests for resolvespec, resolvemcp, funcspec; pgsql adapter preload tests (same connection, error returned); source guard `pkg/common/tx_guard_test.go` (no direct `RunInTransaction`/`BeginTx`, no `Tx = h.db` beyond the allowlisted BeforeHandle placeholders, no pool statements in spec handlers). +- NOTE: resolvespec never fires `AfterRead`/`AfterCreate` (hook types exist, no call site); not tested, pre-existing. +- NOTE: restheadspec total-count cache is process-wide and ignores the record id; read tests call `resetTotalCache`. +- NOTE: `pkg/security` `TestDatabaseAuthenticator` fails with `-count=2` (also on `cd96404`, before this work); use `-count=1`. ## Tests - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/common/adapters/database/pgsql_preload_test.go b/pkg/common/adapters/database/pgsql_preload_test.go index ee3f6e5..6d93103 100644 --- a/pkg/common/adapters/database/pgsql_preload_test.go +++ b/pkg/common/adapters/database/pgsql_preload_test.go @@ -5,6 +5,7 @@ import ( "errors" "strings" "testing" + "time" "github.com/DATA-DOG/go-sqlmock" @@ -54,3 +55,35 @@ func TestSubqueryPreloadErrorIsReturned(t *testing.T) { }) } } + +// With a single pooled connection, a preload that escaped the transaction would +// block on the pool and fail on the context timeout. +func TestSubqueryPreloadRunsOnTheTransaction(t *testing.T) { + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + defer db.Close() + db.SetMaxOpenConns(1) + + mock.ExpectBegin() + mock.ExpectQuery(`FROM parents`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) + mock.ExpectQuery(`FROM children`).WillReturnRows(sqlmock.NewRows([]string{"id", "user_id"}).AddRow(10, 1)) + mock.ExpectCommit() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + var parents []preloadParent + err = NewPgSQLAdapter(db).RunInTransaction(ctx, func(tx common.Database) error { + return tx.NewSelect().Model(&preloadParent{}).PreloadRelation("Children").Scan(ctx, &parents) + }) + if err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if len(parents) != 1 || len(parents[0].Children) != 1 { + t.Fatalf("preloaded children missing: %+v", parents) + } +} diff --git a/pkg/common/tx_guard_test.go b/pkg/common/tx_guard_test.go new file mode 100644 index 0000000..2c943f5 --- /dev/null +++ b/pkg/common/tx_guard_test.go @@ -0,0 +1,99 @@ +package common + +import ( + "os" + "path/filepath" + "regexp" + "strings" + "testing" +) + +// Source-level regression guard for the single-transaction-per-request rule +// (audit/single_tran.md). Runtime tests prove the current paths; this catches a +// new code path that quietly reaches for the pool. + +var guardedSpecs = []string{"resolvespec", "restheadspec", "websocketspec", "mqttspec", "resolvemcp", "funcspec"} + +var ( + directTxRE = regexp.MustCompile(`\.(RunInTransaction|BeginTx)\(`) + poolHookTxRE = regexp.MustCompile(`\bTx(:\s+|\s*=\s*)(h|handler|h\.handler)\.db\b`) + poolQueryRE = regexp.MustCompile(`\b(h|handler)\.db\.(NewSelect|NewInsert|NewUpdate|NewDelete|Exec|Query)\(`) +) + +// allowedPoolHookTx: hook contexts that start life on the pool before the handler +// opens its transaction (BeforeHandle runs before any tx and must be DB-free). +// runInTx replaces Tx with the transaction before any other hook runs. +var allowedPoolHookTx = map[string]int{ + "resolvespec/handler.go": 1, + "websocketspec/handler.go": 1, + "resolvemcp/handler.go": 4, +} + +// allowedPoolQuery: statements outside the request path. +var allowedPoolQuery = map[string]int{ + "resolvemcp/annotation.go": 2, // tool annotations, not a data request +} + +func guardedFiles(t *testing.T) map[string][]string { + t.Helper() + out := map[string][]string{} + for _, spec := range guardedSpecs { + files, err := filepath.Glob(filepath.Join("..", spec, "*.go")) + if err != nil || len(files) == 0 { + t.Fatalf("no sources found for %s: %v", spec, err) + } + for _, f := range files { + if strings.HasSuffix(f, "_test.go") { + continue + } + raw, err := os.ReadFile(f) + if err != nil { + t.Fatal(err) + } + var lines []string + for _, l := range strings.Split(string(raw), "\n") { + if s := strings.TrimSpace(l); strings.HasPrefix(s, "//") { + continue + } + lines = append(lines, l) + } + out[spec+"/"+filepath.Base(f)] = lines + } + } + return out +} + +func countMatches(lines []string, re *regexp.Regexp) int { + n := 0 + for _, l := range lines { + if re.MatchString(l) { + n++ + } + } + return n +} + +func TestNoDirectTransactionsInSpecHandlers(t *testing.T) { + for file, lines := range guardedFiles(t) { + if n := countMatches(lines, directTxRE); n > 0 { + t.Errorf("%s opens a transaction directly (%d): use the handler's runInTx so OnTxBegin fires", file, n) + } + } +} + +func TestHookContextsDoNotRetainThePool(t *testing.T) { + files := guardedFiles(t) + for file, lines := range files { + if got, want := countMatches(lines, poolHookTxRE), allowedPoolHookTx[file]; got != want { + t.Errorf("%s has %d hook contexts set to the pool, allowed %d: hooks must get the transaction", file, got, want) + } + } +} + +func TestSpecHandlersDoNotQueryThePoolDirectly(t *testing.T) { + for file, lines := range guardedFiles(t) { + if got, want := countMatches(lines, poolQueryRE), allowedPoolQuery[file]; got != want { + t.Errorf("%s runs %d statements on the pool, allowed %d: use the transaction", file, got, want) + } + } +} diff --git a/pkg/funcspec/tx_test.go b/pkg/funcspec/tx_test.go index 3ad946f..df95ba2 100644 --- a/pkg/funcspec/tx_test.go +++ b/pkg/funcspec/tx_test.go @@ -9,6 +9,7 @@ import ( "testing" "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/security" ) // txFactory returns a pool whose every transaction is a distinct MockDatabase. @@ -100,3 +101,45 @@ func TestOnTxBeginErrorAnswersTransactionError(t *testing.T) { t.Fatalf("no query may run after a failed OnTxBegin, ran %d", queries) } } + +type stubProvider struct{ security.SecurityProvider } + +func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) { + var execs []string + h := NewHandler(&MockDatabase{ + RunInTransactionFunc: func(ctx context.Context, fn func(common.Database) error) error { + return fn(&MockDatabase{ + ExecFunc: func(ctx context.Context, query string, args ...interface{}) (common.Result, error) { + execs = append(execs, query) + return &MockResult{}, nil + }, + QueryFunc: func(ctx context.Context, dest interface{}, query string, args ...interface{}) error { + execs = append(execs, "QUERY") + if rows, ok := dest.(*[]map[string]interface{}); ok { + *rows = []map[string]interface{}{{"id": float64(1)}} + } + return nil + }, + }) + }, + }) + list, err := security.NewSecurityList(stubProvider{}) + if err != nil { + t.Fatal(err) + } + list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { + return map[string]string{"app.user_id": "1"}, nil + }) + RegisterSecurityHooks(h, list) + + w := httptest.NewRecorder() + h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{})(w, createTestRequest("GET", "/t", nil, nil, nil)) + + if w.Code != http.StatusOK { + t.Fatalf("status %d body %s", w.Code, w.Body) + } + // tx 1: stamp, then the query; tx 2 (BeforeResponse): stamp again. + if len(execs) != 3 || !strings.Contains(execs[0], "set_config('app.user_id'") || execs[1] != "QUERY" || !strings.Contains(execs[2], "set_config('app.user_id'") { + t.Fatalf("each transaction must be stamped before any other SQL, got %v", execs) + } +} diff --git a/pkg/mqttspec/tx_test.go b/pkg/mqttspec/tx_test.go index 4293e92..8ae5acc 100644 --- a/pkg/mqttspec/tx_test.go +++ b/pkg/mqttspec/tx_test.go @@ -89,3 +89,54 @@ func TestHandler_UpdateRunsAfterHookOnSecondTransaction(t *testing.T) { assert.NotEqual(t, begins[0], begins[1]) assert.Equal(t, begins[1], afterTx) } + +func TestHandler_ReadRunsHooksOnOneTransaction(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + + var begins []common.Database + var beforeTx, afterTx common.Database + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { begins = append(begins, c.Tx); return nil }) + handler.hooks.Register(BeforeRead, func(c *HookContext) error { beforeTx = c.Tx; return nil }) + handler.hooks.Register(AfterRead, func(c *HookContext) error { afterTx = c.Tx; return nil }) + + handler.handleRead(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, nil)) + + require.Len(t, begins, 1) + assert.NotEqual(t, handler.db, begins[0]) + assert.Equal(t, begins[0], beforeTx) + assert.Equal(t, begins[0], afterTx) +} + +func TestHandler_CreateRunsHooksOnTwoTransactions(t *testing.T) { + handler, db := setupTestHandler(t) + + var begins []common.Database + var beforeTx, afterTx common.Database + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { begins = append(begins, c.Tx); return nil }) + handler.hooks.Register(BeforeCreate, func(c *HookContext) error { beforeTx = c.Tx; return nil }) + handler.hooks.Register(AfterCreate, func(c *HookContext) error { afterTx = c.Tx; return nil }) + + hook := newTxHook(handler, map[string]interface{}{"id": 5, "name": "n", "email": "n@example.com", "status": "active"}) + hook.ID = "" + handler.handleCreate(&Client{ID: "c1"}, &Message{ID: "m1"}, hook) + + var got TestUser + require.NoError(t, db.First(&got, 5).Error) + require.Len(t, begins, 2) + assert.NotEqual(t, begins[0], begins[1]) + assert.Equal(t, begins[0], beforeTx) + assert.Equal(t, begins[1], afterTx) +} + +func TestHandler_BeforeHookErrorAbortsUpdateWithoutWriting(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + handler.hooks.Register(BeforeUpdate, func(c *HookContext) error { return errors.New("denied") }) + + handler.handleUpdate(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, map[string]interface{}{"name": "b"})) + + var got TestUser + require.NoError(t, db.First(&got, 1).Error) + assert.Equal(t, "a", got.Name) +} diff --git a/pkg/resolvemcp/tx_test.go b/pkg/resolvemcp/tx_test.go index 80f4945..3bff822 100644 --- a/pkg/resolvemcp/tx_test.go +++ b/pkg/resolvemcp/tx_test.go @@ -11,6 +11,7 @@ import ( "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" "github.com/bitechdev/ResolveSpec/pkg/modelregistry" + "github.com/bitechdev/ResolveSpec/pkg/security" ) type txItem struct { @@ -205,3 +206,121 @@ func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) { } } } + +func TestAfterDeleteErrorRollsBackDelete(t *testing.T) { + h, mock, ctx := newTxHarness(t) + h.Hooks().Register(AfterDelete, func(*HookContext) error { return sql.ErrConnDone }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + if _, err := h.executeDelete(ctx, "public", "items", "7"); err == nil { + t.Fatal("a failing AfterDelete must fail the request") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestCreateBeforeHookErrorRollsBackWithoutInsert(t *testing.T) { + h, mock, ctx := newTxHarness(t) + h.Hooks().Register(BeforeCreate, func(*HookContext) error { return sql.ErrConnDone }) + + mock.ExpectBegin() + mock.ExpectRollback() + + if _, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a"}); err == nil { + t.Fatal("expected error") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestCreateInsertErrorRollsBackBatch(t *testing.T) { + h, mock, ctx := newTxHarness(t) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) + mock.ExpectQuery(`INSERT`).WillReturnError(sql.ErrConnDone) + mock.ExpectRollback() + + items := []interface{}{map[string]interface{}{"name": "a"}, map[string]interface{}{"name": "b"}} + if _, err := h.executeCreate(ctx, "public", "items", items); err == nil { + t.Fatal("expected error") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestOnTxBeginErrorRollsBackEveryOperation(t *testing.T) { + ops := map[string]func(h *Handler, ctx context.Context) error{ + "read": func(h *Handler, ctx context.Context) error { + _, _, err := h.executeRead(ctx, "public", "items", "7", common.RequestOptions{}) + return err + }, + "create": func(h *Handler, ctx context.Context) error { + _, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a"}) + return err + }, + "update": func(h *Handler, ctx context.Context) error { + _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "a"}) + return err + }, + "delete": func(h *Handler, ctx context.Context) error { + _, err := h.executeDelete(ctx, "public", "items", "7") + return err + }, + } + for name, op := range ops { + t.Run(name, func(t *testing.T) { + h, mock, ctx := newTxHarness(t) + h.Hooks().Register(OnTxBegin, func(*HookContext) error { return sql.ErrConnDone }) + mock.ExpectBegin() + mock.ExpectRollback() + if err := op(h, ctx); err == nil { + t.Fatal("expected error") + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + }) + } +} + +type stubProvider struct{ security.SecurityProvider } + +func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) { + h, mock, ctx := newTxHarness(t) + list, err := security.NewSecurityList(stubProvider{}) + if err != nil { + t.Fatal(err) + } + list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { + return map[string]string{"app.user_id": "7"}, nil + }) + RegisterSecurityHooks(h, list) + ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + + // Update opens two transactions; each must be stamped before any other SQL. + cols := []string{"id", "name"} + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b")) + mock.ExpectCommit() + + if _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "b"}); err != nil { + t.Fatal(err) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/resolvespec/ops_tx_test.go b/pkg/resolvespec/ops_tx_test.go new file mode 100644 index 0000000..aa75cba --- /dev/null +++ b/pkg/resolvespec/ops_tx_test.go @@ -0,0 +1,186 @@ +package resolvespec + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func opCtx(t *testing.T) (context.Context, common.ResponseWriter, *httptest.ResponseRecorder) { + t.Helper() + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPost, "/", nil)) + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + t.Cleanup(cancel) + return WithRequestData(base, "public", "items", "items", &delItem{}, &delItem{}), w, rec +} + +// hookTrace records hook order and the Tx each hook saw. +type hookTrace struct { + order []string + tx map[HookType][]common.Database +} + +func traceHooks(h *Handler, types ...HookType) *hookTrace { + tr := &hookTrace{tx: map[HookType][]common.Database{}} + for _, ht := range types { + ht := ht + h.Hooks().Register(ht, func(c *HookContext) error { + tr.order = append(tr.order, string(ht)) + tr.tx[ht] = append(tr.tx[ht], c.Tx) + return nil + }) + } + return tr +} + +func (tr *hookTrace) mustBeOn(t *testing.T, tx common.Database, types ...HookType) { + t.Helper() + for _, ht := range types { + if len(tr.tx[ht]) == 0 || tr.tx[ht][0] != tx { + t.Fatalf("%s must run on the OnTxBegin transaction", ht) + } + } +} + +func TestReadRunsHooksOnOneTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeRead) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + h.handleRead(ctx, w, "7", 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(tr.tx[OnTxBegin]) != 1 || tr.order[0] != "on_tx_begin" { + t.Fatalf("OnTxBegin must fire once and first, got %v", tr.order) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], BeforeRead) +} + +func TestReadBeforeHookErrorRollsBackWithoutQueries(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(BeforeRead, func(*HookContext) error { return errors.New("denied") }) + + mock.ExpectBegin() + mock.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleRead(ctx, w, "7", common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing BeforeRead must not return data: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestCreateRunsHooksOnTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeCreate) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + h.handleCreate(ctx, w, map[string]interface{}{"name": "a"}, 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 tr.order[0] != "on_tx_begin" { + t.Fatalf("OnTxBegin must fire first, got %v", tr.order) + } + for _, tx := range tr.tx[BeforeCreate] { + if tx == nil || tx == h.db { + t.Fatal("BeforeCreate must not get the pool") + } + } +} + +func TestCreateBeforeHookErrorRollsBack(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(BeforeCreate, func(*HookContext) error { return errors.New("denied") }) + + mock.ExpectBegin() + mock.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleCreate(ctx, w, map[string]interface{}{"name": "a"}, common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing BeforeCreate must not create: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestUpdateAfterHookErrorRollsBackAndSkipsRefetch(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(AfterUpdate, func(*HookContext) error { return errors.New("audit failed") }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleUpdate(ctx, w, "7", nil, map[string]interface{}{"name": "b"}, common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing AfterUpdate must fail the request: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestUpdateRunsBeforeAndAfterOnFirstTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, BeforeUpdate, AfterUpdate) + + 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() + + ctx, w, rec := opCtx(t) + 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(tr.tx[OnTxBegin]) != 2 { + t.Fatalf("expected OnTxBegin twice, got %d", len(tr.tx[OnTxBegin])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], BeforeUpdate, AfterUpdate) +} diff --git a/pkg/restheadspec/read_tx_test.go b/pkg/restheadspec/read_tx_test.go index 22ed88b..ac05387 100644 --- a/pkg/restheadspec/read_tx_test.go +++ b/pkg/restheadspec/read_tx_test.go @@ -11,12 +11,14 @@ import ( "github.com/uptrace/bun" "github.com/uptrace/bun/dialect/pgdialect" + "github.com/bitechdev/ResolveSpec/pkg/cache" "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" "github.com/bitechdev/ResolveSpec/pkg/modelregistry" ) func TestAfterReadRunsInSecondTransaction(t *testing.T) { + resetTotalCache(t) sqlDB, mock, err := sqlmock.New() if err != nil { t.Fatal(err) @@ -60,3 +62,100 @@ func TestAfterReadRunsInSecondTransaction(t *testing.T) { t.Fatal("BeforeRead must run on the first tx and AfterRead on the second") } } + +// resetTotalCache empties the process-wide query-total cache. Its key ignores the +// record id, so a cached total would skip the count query and desync the mock. +func resetTotalCache(t *testing.T) { + t.Helper() + _ = cache.GetDefaultCache().Clear(context.Background()) + t.Cleanup(func() { _ = cache.GetDefaultCache().Clear(context.Background()) }) +} + +func newBunHarness(t *testing.T) (*Handler, sqlmock.Sqlmock) { + t.Helper() + sqlDB, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + sqlDB.SetMaxOpenConns(1) + t.Cleanup(func() { _ = sqlDB.Close() }) + return NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry()), mock +} + +func itemCtx(t *testing.T) context.Context { + t.Helper() + base, cancel := context.WithTimeout(context.Background(), 2*time.Second) + t.Cleanup(cancel) + ctx := WithSchema(base, "public") + ctx = WithEntity(ctx, "items") + ctx = WithTableName(ctx, "items") + return WithModel(ctx, delItem{}) +} + +func TestAfterReadErrorFailsRequestOnSecondTransaction(t *testing.T) { + resetTotalCache(t) + h, mock := newBunHarness(t) + h.Hooks().Register(AfterRead, func(*HookContext) error { return http.ErrAbortHandler }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectRollback() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + h.handleRead(itemCtx(t), w, "7", ExtendedRequestOptions{}) + + if rec.Code != http.StatusInternalServerError { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestBeforeReadErrorRollsBackWithoutQueries(t *testing.T) { + h, mock := newBunHarness(t) + h.Hooks().Register(BeforeRead, func(*HookContext) error { return http.ErrAbortHandler }) + + mock.ExpectBegin() + mock.ExpectRollback() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + h.handleRead(itemCtx(t), w, "7", ExtendedRequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing BeforeRead must not return data: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestAfterUpdateErrorFailsRequest(t *testing.T) { + h, mock := newBunHarness(t) + h.Hooks().Register(AfterUpdate, func(*HookContext) error { return http.ErrAbortHandler }) + + 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.ExpectRollback() + + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPut, "/", nil)) + h.handleUpdate(itemCtx(t), w, "7", nil, map[string]interface{}{"name": "b"}, ExtendedRequestOptions{}) + + if rec.Code != http.StatusInternalServerError { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/websocketspec/tx_test.go b/pkg/websocketspec/tx_test.go index 12e410d..979ec7e 100644 --- a/pkg/websocketspec/tx_test.go +++ b/pkg/websocketspec/tx_test.go @@ -8,6 +8,8 @@ import ( "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" @@ -171,3 +173,79 @@ func TestOnTxBeginErrorRollsBackWithoutDetail(t *testing.T) { t.Fatalf("unexpected response %+v", resp) } } + +func TestCreateRunsHooksOnTwoTransactions(t *testing.T) { + // The bun adapter builds model-based inserts; the pgsql adapter does not. + h, mock, conn, hookCtx := newTxHarness(t) + sqlDB, bunMock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + sqlDB.SetMaxOpenConns(1) + t.Cleanup(func() { _ = sqlDB.Close() }) + h.db = database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())) + conn = NewConnection("c2", nil, h) + mock = bunMock + hookCtx.ID = "" + hookCtx.Data = map[string]interface{}{"id": 7, "name": "a"} + var begins []common.Database + var beforeTx, afterTx common.Database + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { begins = append(begins, c.Tx); return nil }) + h.Hooks().Register(BeforeCreate, func(c *HookContext) error { beforeTx = c.Tx; return nil }) + h.Hooks().Register(AfterCreate, func(c *HookContext) error { afterTx = c.Tx; return nil }) + + mock.ExpectBegin() + mock.ExpectExec(`INSERT`).WillReturnResult(sqlmock.NewResult(7, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + h.handleCreate(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if !response(t, conn).Success { + t.Fatal("expected success") + } + if len(begins) != 2 || begins[0] == begins[1] || beforeTx != begins[0] || afterTx != begins[1] { + t.Fatalf("BeforeCreate must run on tx 1 and AfterCreate on tx 2, begins=%d", len(begins)) + } +} + +func TestBeforeHookErrorRollsBackWithoutWriting(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + hookCtx.Data = map[string]interface{}{"name": "b"} + h.Hooks().Register(BeforeUpdate, func(*HookContext) error { return errors.New("denied") }) + + mock.ExpectBegin() + mock.ExpectRollback() + + h.handleUpdate(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if resp := response(t, conn); resp.Success || resp.Error == nil || resp.Error.Code != "hook_error" { + t.Fatalf("unexpected response %+v", resp) + } +} + +func TestAfterDeleteErrorRollsBackDelete(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + h.Hooks().Register(AfterDelete, func(*HookContext) error { return errors.New("audit failed") }) + + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + h.handleDelete(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if response(t, conn).Success { + t.Fatal("a failing AfterDelete must fail the request") + } +}