mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 12:31:59 +00:00
fix: fire resolvespec AfterRead/AfterCreate/AfterDelete, key restheadspec total cache by record id
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user