From a65ca5f5ce877514c5faeb6855a58a294da610bf Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 30 Sep 2026 23:40:32 +0200 Subject: [PATCH] fix: fire resolvespec AfterRead/AfterCreate/AfterDelete, key restheadspec total cache by record id --- audit/single_tran.md | 5 +- pkg/common/TRANSACTIONS.md | 2 + pkg/common/tx_guard_test.go | 45 ++++++ pkg/resolvespec/handler.go | 98 +++++++++++- pkg/resolvespec/ops_tx_test.go | 237 +++++++++++++++++++++++++++++ pkg/restheadspec/cache_helpers.go | 12 +- pkg/restheadspec/cache_key_test.go | 61 ++++++++ pkg/restheadspec/handler.go | 1 + 8 files changed, 448 insertions(+), 13 deletions(-) create mode 100644 pkg/restheadspec/cache_key_test.go diff --git a/audit/single_tran.md b/audit/single_tran.md index 380f30e..13bb6a6 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -97,8 +97,8 @@ - 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. - 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`. +- FIXED: resolvespec now fires `AfterRead` (in the read tx; `Result` = scanned slice for single and list) `AfterCreate` (in the create tx, per record, all four create paths) and `AfterDelete` (in the delete tx, once per request; a failure rolls the delete back). Before, neither fired, so `AfterRead` column-level security masking (registered by `RegisterSecurityHooks`) was silently skipped on resolvespec reads. Failing `AfterRead`/`AfterCreate` fails the request and rolls back. Tests: `pkg/resolvespec/ops_tx_test.go` (incl. end-to-end column hiding). +- FIXED: restheadspec total-count cache key now includes the record id; a read by id (total 1) no longer poisons the list total for the 2-minute TTL. Tests: `pkg/restheadspec/cache_key_test.go`. resolvespec is unaffected (its count runs before the id filter, so the total is the list total by design). The cache is still process-wide, so read tests call `resetTotalCache`. - NOTE: `pkg/security` `TestDatabaseAuthenticator` fails with `-count=2` (also on `cd96404`, before this work); use `-count=1`. ## Tests @@ -118,3 +118,4 @@ - `dbtrace` shows `pooled=0` for every handler op on a hooked model. - RLS GUC set in `OnTxBegin` is visible to read, create, update, delete queries and hooks. - No `Tx: h.db` / `hookCtx.Tx = h.db` left in spec handlers. +- OPEN: websocketspec `BeforeDisconnect`/`AfterDisconnect` are defined but never executed (connection lifecycle, not DB). Allowlisted in `TestEveryDefinedHookHasACallSite`; wire them to remove the entry. diff --git a/pkg/common/TRANSACTIONS.md b/pkg/common/TRANSACTIONS.md index df2f8d2..5b72d59 100644 --- a/pkg/common/TRANSACTIONS.md +++ b/pkg/common/TRANSACTIONS.md @@ -21,6 +21,8 @@ Every DB statement and every DB-touching hook of one request runs on one transac \* resolvespec, websocketspec, mqttspec, resolvemcp. restheadspec runs `AfterRead` in tx 2. † restheadspec, websocketspec, mqttspec run `AfterUpdate` in tx 2. resolvespec, resolvemcp run it in tx 1. +resolvespec has no tx 2 for create: `AfterCreate` runs in tx 1, once per record (per item in a batch), after the insert and re-fetch. Its `AfterRead` gets the scanned slice (single and list reads) and may mask in place; a failing `AfterRead` fails the read (fail closed). + Tx 2 exists so the re-fetch sees trigger changes from the committed write. ## Per spec diff --git a/pkg/common/tx_guard_test.go b/pkg/common/tx_guard_test.go index 2c943f5..40a3820 100644 --- a/pkg/common/tx_guard_test.go +++ b/pkg/common/tx_guard_test.go @@ -97,3 +97,48 @@ func TestSpecHandlersDoNotQueryThePoolDirectly(t *testing.T) { } } } + +// Hook types that are defined but deliberately or knowingly never executed. +// Anything else defined in a spec's hooks.go must have an Execute call site: an +// unwired hook silently disables whatever is registered on it (resolvespec's +// AfterRead skipped column-level security masking until it was wired). +var unwiredHooks = map[string]string{ + "websocketspec/BeforeDisconnect": "connection close is not hooked yet", + "websocketspec/AfterDisconnect": "connection close is not hooked yet", +} + +var hookConstRE = regexp.MustCompile(`(?m)^\s*([A-Z][A-Za-z0-9]*)\s+HookType\s*=`) + +func TestEveryDefinedHookHasACallSite(t *testing.T) { + for _, spec := range []string{"resolvespec", "restheadspec", "websocketspec", "resolvemcp", "funcspec"} { + raw, err := os.ReadFile(filepath.Join("..", spec, "hooks.go")) + if err != nil { + t.Fatal(err) + } + var src strings.Builder + files, _ := filepath.Glob(filepath.Join("..", spec, "*.go")) + for _, f := range files { + if strings.HasSuffix(f, "_test.go") || strings.HasSuffix(f, "hooks.go") || strings.HasSuffix(f, "hooks_example.go") { + continue + } + b, err := os.ReadFile(f) + if err != nil { + t.Fatal(err) + } + src.Write(b) + } + for _, m := range hookConstRE.FindAllStringSubmatch(string(raw), -1) { + name := m[1] + if name == "BeforeOp" || name == "OnTxBegin" { // fired by the registry / runInTx + continue + } + if _, ok := unwiredHooks[spec+"/"+name]; ok { + continue + } + call := regexp.MustCompile(`Execute(BeforeOp)?\(` + name + `\b`) + if !call.MatchString(src.String()) { + t.Errorf("%s: hook %s is defined but never executed", spec, name) + } + } + } +} diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index 8827a15..da30452 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -666,6 +666,17 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st result = reflect.ValueOf(modelPtr).Elem().Interface() } + // AfterRead runs inside the read transaction (e.g. column-level security + // masking). Result is the scanned slice for single and multi-record reads + // alike; hooks mutate the records in place, which `result` shares. + hookCtx.Result = modelPtr + hookCtx.Error = nil + if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { + logger.Error("AfterRead hook failed: %v", err) + statusCode, errCode, errMsg = http.StatusInternalServerError, "hook_error", "Hook execution failed" + return err + } + logger.Info("Successfully retrieved records") return nil }) @@ -762,7 +773,17 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat var procErr error nestedResult, procErr = h.nestedProcessor.ProcessNestedCUD(ctx, "insert", v, model, make(map[string]interface{}), tableName) - return procErr + if procErr != nil { + return procErr + } + res, err := h.afterCreate(hookCtx, nestedResult.Data) + if err != nil { + return err + } + if m, ok := res.(map[string]interface{}); ok { + nestedResult.Data = m + } + return nil }) if err != nil { logger.Error("Error in nested create: %v", err) @@ -814,6 +835,11 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat return err } logger.Info("Successfully created record, rows affected: %d", result.RowsAffected()) + res, err := h.afterCreate(hookCtx, responseData) + if err != nil { + return err + } + responseData = res return nil } @@ -830,6 +856,11 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat } else { logger.Warn("Failed to re-fetch created record with %s=%v: %v", pkName, insertedID, fetchErr) } + res, err := h.afterCreate(hookCtx, responseData) + if err != nil { + return err + } + responseData = res return nil }) if err != nil { @@ -889,6 +920,13 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if err != nil { return fmt.Errorf("failed to process item: %w", err) } + res, err := h.afterCreate(hookCtx, result.Data) + if err != nil { + return err + } + if m, ok := res.(map[string]interface{}); ok { + result.Data = m + } results = append(results, result.Data) } return nil @@ -941,7 +979,11 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if _, err := txQuery.Exec(ctx); err != nil { return err } - responseItems = append(responseItems, item) + res, err := h.afterCreate(hookCtx, item) + if err != nil { + return err + } + responseItems = append(responseItems, res) continue } var returnedID interface{} @@ -949,14 +991,20 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat return err } fetchedRecord := reflect.New(modelElemType).Interface() + var created interface{} if fetchErr := tx.NewSelect().Model(fetchedRecord). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), returnedID). ScanModel(ctx); fetchErr == nil { - responseItems = append(responseItems, mergeWithInput(fetchedRecord, item)) + created = mergeWithInput(fetchedRecord, item) } else { logger.Warn("Failed to re-fetch created record with %s=%v: %v", pkName, returnedID, fetchErr) - responseItems = append(responseItems, item) + created = item } + res, err := h.afterCreate(hookCtx, created) + if err != nil { + return err + } + responseItems = append(responseItems, res) } return nil }) @@ -1022,6 +1070,13 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if err != nil { return fmt.Errorf("failed to process item: %w", err) } + res, err := h.afterCreate(hookCtx, result.Data) + if err != nil { + return err + } + if m, ok := res.(map[string]interface{}); ok { + result.Data = m + } results = append(results, result.Data) } } @@ -1080,7 +1135,11 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if _, err := txQuery.Exec(ctx); err != nil { return err } - responseItems = append(responseItems, itemMap) + res, err := h.afterCreate(hookCtx, itemMap) + if err != nil { + return err + } + responseItems = append(responseItems, res) continue } var returnedID interface{} @@ -1088,14 +1147,20 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat return err } fetchedRecord := reflect.New(modelElemType).Interface() + var created interface{} if fetchErr := tx.NewSelect().Model(fetchedRecord). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), returnedID). ScanModel(ctx); fetchErr == nil { - responseItems = append(responseItems, mergeWithInput(fetchedRecord, itemMap)) + created = mergeWithInput(fetchedRecord, itemMap) } else { logger.Warn("Failed to re-fetch created record with %s=%v: %v", pkName, returnedID, fetchErr) - responseItems = append(responseItems, itemMap) + created = itemMap } + res, err := h.afterCreate(hookCtx, created) + if err != nil { + return err + } + responseItems = append(responseItems, res) } return nil }) @@ -1677,6 +1742,15 @@ func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id if failure != nil { return failure } + // AfterDelete runs inside the delete transaction: a failing hook rolls the + // delete back. Result is the deleted record, or the batch summary. + hookCtx.Result = payload + if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { + logger.Error("AfterDelete hook failed: %v", err) + failure = &deleteFailure{http.StatusInternalServerError, "hook_error", "Hook execution failed", err} + return failure + } + payload = hookCtx.Result return nil }) if failure != nil { @@ -2598,3 +2672,13 @@ func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(t return h.hooks.Execute(OnTxBegin, hookCtx) }, body) } + +// afterCreate runs the AfterCreate hooks inside the create transaction with the +// created record as Result and returns the (possibly replaced) result. +func (h *Handler) afterCreate(hookCtx *HookContext, result interface{}) (interface{}, error) { + hookCtx.Result = result + if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { + return nil, fmt.Errorf("AfterCreate hook failed: %w", err) + } + return hookCtx.Result, nil +} diff --git a/pkg/resolvespec/ops_tx_test.go b/pkg/resolvespec/ops_tx_test.go index aa75cba..b5da745 100644 --- a/pkg/resolvespec/ops_tx_test.go +++ b/pkg/resolvespec/ops_tx_test.go @@ -3,16 +3,28 @@ package resolvespec import ( "context" "errors" + "fmt" "net/http" "net/http/httptest" + "strings" "testing" "time" "github.com/DATA-DOG/go-sqlmock" + "github.com/bitechdev/ResolveSpec/pkg/cache" "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/security" ) +// resetTotalCache empties the process-wide query-total cache so a total cached by +// another test cannot 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 opCtx(t *testing.T) (context.Context, common.ResponseWriter, *httptest.ResponseRecorder) { t.Helper() rec := httptest.NewRecorder() @@ -51,6 +63,7 @@ func (tr *hookTrace) mustBeOn(t *testing.T, tx common.Database, types ...HookTyp } func TestReadRunsHooksOnOneTransaction(t *testing.T) { + resetTotalCache(t) h, mock, _ := newDeleteHarness(t) tr := traceHooks(h, OnTxBegin, BeforeRead) @@ -184,3 +197,227 @@ func TestUpdateRunsBeforeAndAfterOnFirstTransaction(t *testing.T) { } tr.mustBeOn(t, tr.tx[OnTxBegin][0], BeforeUpdate, AfterUpdate) } + +func TestAfterReadRunsOnReadTransactionForSingleAndList(t *testing.T) { + for name, id := range map[string]string{"single": "7", "list": ""} { + t.Run(name, func(t *testing.T) { + resetTotalCache(t) + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, AfterRead) + var resultType string + h.Hooks().Register(AfterRead, func(c *HookContext) error { + resultType = fmt.Sprintf("%T", c.Result) + return nil + }) + + 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, id, 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[AfterRead]) != 1 { + t.Fatalf("AfterRead must fire once, got %d", len(tr.tx[AfterRead])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], AfterRead) + if !strings.HasPrefix(resultType, "*[]") { + t.Fatalf("AfterRead Result must be the scanned slice, got %s", resultType) + } + }) + } +} + +func TestAfterReadErrorFailsReadAndRollsBack(t *testing.T) { + resetTotalCache(t) + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(AfterRead, func(*HookContext) error { return errors.New("masking failed") }) + + 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.ExpectRollback() + + ctx, w, rec := opCtx(t) + h.handleRead(ctx, w, "7", common.RequestOptions{}) + + if rec.Code == http.StatusOK { + t.Fatalf("a failing AfterRead must not return data (fail closed): %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +type columnSecProvider struct{ security.SecurityProvider } + +func (columnSecProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]security.ColumnSecurity, error) { + return []security.ColumnSecurity{{ + Schema: schema, Tablename: table, Path: []string{"name"}, UserID: userID, Accesstype: "hide", + }}, nil +} + +func (columnSecProvider) GetRowSecurity(ctx context.Context, userRef any, schema, table string) (security.RowSecurity, error) { + return security.RowSecurity{}, nil +} + +// Column-level security is applied by an AfterRead hook; before AfterRead was wired +// into resolvespec reads, the configured masking was silently skipped. +func TestColumnSecurityHidesColumnOnRead(t *testing.T) { + for name, id := range map[string]string{"single": "7", "list": ""} { + t.Run(name, func(t *testing.T) { + resetTotalCache(t) + h, mock, _ := newDeleteHarness(t) + list, err := security.NewSecurityList(columnSecProvider{}) + if err != nil { + t.Fatal(err) + } + RegisterSecurityHooks(h, list) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "secret")) + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"}) + ctx = context.WithValue(ctx, security.UserIDKey, 7) + h.handleRead(ctx, w, id, common.RequestOptions{}) + + if rec.Code != http.StatusOK { + t.Fatalf("status %d body %s", rec.Code, rec.Body) + } + if strings.Contains(rec.Body.String(), "secret") { + t.Fatalf("hidden column leaked: %s", rec.Body) + } + }) + } +} + +func TestAfterCreateRunsInsideCreateTransaction(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, AfterCreate) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + 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 len(tr.tx[AfterCreate]) != 1 { + t.Fatalf("AfterCreate must fire once, got %d", len(tr.tx[AfterCreate])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], AfterCreate) +} + +func TestAfterCreateFiresPerItemInBatch(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, AfterCreate) + + mock.ExpectBegin() + for i := 1; i <= 2; i++ { + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(i)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(i, "a")) + } + mock.ExpectCommit() + + ctx, w, rec := opCtx(t) + items := []interface{}{map[string]interface{}{"name": "a"}, map[string]interface{}{"name": "b"}} + h.handleCreate(ctx, w, items, 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[AfterCreate]) != 2 || len(tr.tx[OnTxBegin]) != 1 { + t.Fatalf("AfterCreate must fire per item in one transaction, got %d in %d tx", len(tr.tx[AfterCreate]), len(tr.tx[OnTxBegin])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], AfterCreate) +} + +func TestAfterCreateErrorRollsBackCreate(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(AfterCreate, func(*HookContext) error { return errors.New("audit failed") }) + + mock.ExpectBegin() + mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + 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 AfterCreate must fail the request: %s", rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} + +func TestAfterDeleteRunsInsideDeleteTransaction(t *testing.T) { + for name, tc := range map[string]struct { + id string + data interface{} + exec int + }{"single": {id: "7", exec: 1}, "batch": {data: []interface{}{"1", "2"}, exec: 2}} { + t.Run(name, func(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + tr := traceHooks(h, OnTxBegin, AfterDelete) + + 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(tr.tx[AfterDelete]) != 1 { + t.Fatalf("AfterDelete must fire once per request, got %d", len(tr.tx[AfterDelete])) + } + tr.mustBeOn(t, tr.tx[OnTxBegin][0], AfterDelete) + }) + } +} + +func TestAfterDeleteErrorRollsBackDelete(t *testing.T) { + h, mock, _ := newDeleteHarness(t) + h.Hooks().Register(AfterDelete, func(*HookContext) error { return errors.New("audit failed") }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectRollback() + + if rec := runDelete(h, "7", nil); rec.Code != http.StatusInternalServerError { + t.Fatalf("a failing AfterDelete must fail and roll back the delete: status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } +} diff --git a/pkg/restheadspec/cache_helpers.go b/pkg/restheadspec/cache_helpers.go index 1e81187..c0950b9 100644 --- a/pkg/restheadspec/cache_helpers.go +++ b/pkg/restheadspec/cache_helpers.go @@ -22,6 +22,7 @@ type expandOptionKey struct { // queryCacheKey represents the components used to build a cache key for query total count type queryCacheKey struct { TableName string `json:"table_name"` + ID string `json:"id,omitempty"` Filters []common.FilterOption `json:"filters"` Sort []common.SortOption `json:"sort"` CustomSQLWhere string `json:"custom_sql_where,omitempty"` @@ -39,12 +40,15 @@ type cachedTotal struct { } // buildExtendedQueryCacheKey builds a cache key for extended query options (restheadspec) -// Includes expand, distinct, and cursor pagination options -func buildExtendedQueryCacheKey(tableName string, filters []common.FilterOption, sort []common.SortOption, +// Includes expand, distinct, and cursor pagination options. id is the record id of a +// single-record read: it constrains the counted query, so a count cached for one id +// (or for the whole list) must never be served for another. +func buildExtendedQueryCacheKey(tableName, id string, filters []common.FilterOption, sort []common.SortOption, customWhere, customOr string, customJoin []string, expandOpts []interface{}, distinct bool, cursorFwd, cursorBwd string) string { key := queryCacheKey{ TableName: tableName, + ID: id, Filters: filters, Sort: sort, CustomSQLWhere: customWhere, @@ -77,8 +81,8 @@ func buildExtendedQueryCacheKey(tableName string, filters []common.FilterOption, jsonData, err := json.Marshal(key) if err != nil { // Fallback to simple string concatenation if JSON fails - return hashString(fmt.Sprintf("%s_%v_%v_%s_%s_%v_%v_%v_%s_%s", - tableName, filters, sort, customWhere, customOr, customJoin, expandOpts, distinct, cursorFwd, cursorBwd)) + return hashString(fmt.Sprintf("%s_%s_%v_%v_%s_%s_%v_%v_%v_%s_%s", + tableName, id, filters, sort, customWhere, customOr, customJoin, expandOpts, distinct, cursorFwd, cursorBwd)) } return hashString(string(jsonData)) diff --git a/pkg/restheadspec/cache_key_test.go b/pkg/restheadspec/cache_key_test.go new file mode 100644 index 0000000..cd396ff --- /dev/null +++ b/pkg/restheadspec/cache_key_test.go @@ -0,0 +1,61 @@ +package restheadspec + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func TestQueryTotalCacheKeyIncludesRecordID(t *testing.T) { + list := buildExtendedQueryCacheKey("items", "", nil, nil, "", "", nil, nil, false, "", "") + one := buildExtendedQueryCacheKey("items", "7", nil, nil, "", "", nil, nil, false, "", "") + other := buildExtendedQueryCacheKey("items", "8", nil, nil, "", "", nil, nil, false, "", "") + if list == one || one == other { + t.Fatal("list, id 7 and id 8 must not share a cache key") + } + if one != buildExtendedQueryCacheKey("items", "7", nil, nil, "", "", nil, nil, false, "", "") { + t.Fatal("the key must be stable for the same query") + } +} + +// A read by id counts 1 row; before the id was part of the key, that total was +// cached under the list query's key and served as the list total for 2 minutes. +func TestReadByIDDoesNotPoisonListTotal(t *testing.T) { + resetTotalCache(t) + h, mock := newBunHarness(t) + cols := []string{"id", "name"} + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a")) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectCommit() + // The list must run its own count query, not reuse the id read's total. + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(5)) + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a").AddRow(8, "b")) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectCommit() + + read := func(id string) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + h.handleRead(itemCtx(t), w, id, ExtendedRequestOptions{}) + return rec + } + if rec := read("7"); rec.Code != http.StatusOK { + t.Fatalf("read by id: status %d body %s", rec.Code, rec.Body) + } + if rec := read(""); rec.Code != http.StatusOK { + t.Fatalf("list: status %d body %s", rec.Code, rec.Body) + } + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatalf("the list must run its own count query: %v", err) + } +} diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index 649901f..29a59ac 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -824,6 +824,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st cacheKeyHash := buildExtendedQueryCacheKey( tableName, + id, options.Filters, options.Sort, options.CustomSQLWhere,