feat(tx): run resolvemcp operations in per-operation transactions with OnTxBegin

This commit is contained in:
2026-09-30 22:49:47 +02:00
parent ed457eb14a
commit 4cbe4f597d
4 changed files with 350 additions and 107 deletions
+131 -105
View File
@@ -247,10 +247,27 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
return nil, nil, err
}
// Hooks and queries share one transaction so transaction-local state set by
// hooks (e.g. RLS settings) applies to every statement.
var data interface{}
var metadata *common.Metadata
err = h.runInTx(ctx, hookCtx, func(common.Database) error {
var err error
data, metadata, err = h.readInTx(ctx, hookCtx, model, modelType, tableName, id, options)
return err
})
if err != nil {
return nil, nil, err
}
return data, metadata, nil
}
// readInTx runs the read hooks and queries on hookCtx.Tx.
func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model interface{}, modelType reflect.Type, tableName, id string, options common.RequestOptions) (interface{}, *common.Metadata, error) {
sliceType := reflect.SliceOf(reflect.PointerTo(modelType))
modelPtr := reflect.New(sliceType).Interface()
query := h.db.NewSelect().Model(modelPtr)
query := hookCtx.Tx.NewSelect().Model(modelPtr)
tempInstance := reflect.New(modelType).Interface()
if provider, ok := tempInstance.(common.TableNameProvider); !ok || provider.TableName() == "" {
@@ -431,96 +448,83 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
return nil, err
}
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
return nil, err
}
// Use potentially modified data
data = hookCtx.Data
pkName := reflection.GetPrimaryKeyName(model)
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
modelType = modelType.Elem()
}
switch v := data.(type) {
case map[string]interface{}:
query := h.db.NewInsert().Table(tableName)
for key, value := range v {
query = query.Value(key, value)
// Transaction 1: BeforeCreate + inserts.
var (
single bool
originals []map[string]interface{}
insertedIDs []interface{}
)
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
return err
}
if pkName != "" {
var insertedID interface{}
if err := query.Returning(pkName).Scan(ctx, &insertedID); err != nil {
return nil, fmt.Errorf("create error: %w", err)
}
// Re-fetch after insert to capture DB-generated defaults/triggers.
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
modelType = modelType.Elem()
}
fetchedRecord := reflect.New(modelType).Interface()
if err := h.db.NewSelect().Model(fetchedRecord).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), insertedID).
ScanModel(ctx); err == nil {
v = mergeWithInput(fetchedRecord, v)
} else {
logger.Warn("Failed to re-fetch created record with %s=%v: %v", pkName, insertedID, err)
}
} else {
if _, err := query.Exec(ctx); err != nil {
return nil, fmt.Errorf("create error: %w", err)
}
}
hookCtx.Result = v
if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil {
return nil, fmt.Errorf("AfterCreate hook failed: %w", err)
}
return v, nil
case []interface{}:
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
modelType = modelType.Elem()
}
originals := make([]map[string]interface{}, 0, len(v))
insertedIDs := make([]interface{}, 0, len(v))
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
// Use potentially modified data
switch v := hookCtx.Data.(type) {
case map[string]interface{}:
single = true
originals = []map[string]interface{}{v}
case []interface{}:
originals = make([]map[string]interface{}, 0, len(v))
for _, item := range v {
itemMap, ok := item.(map[string]interface{})
if !ok {
return fmt.Errorf("each item must be an object")
}
q := tx.NewInsert().Table(tableName)
for key, value := range itemMap {
q = q.Value(key, value)
}
if pkName == "" {
if _, err := q.Exec(ctx); err != nil {
return err
}
originals = append(originals, itemMap)
insertedIDs = append(insertedIDs, nil)
continue
}
var returnedID interface{}
if err := q.Returning(pkName).Scan(ctx, &returnedID); err != nil {
originals = append(originals, itemMap)
}
default:
return fmt.Errorf("data must be an object or array of objects")
}
insertedIDs = make([]interface{}, 0, len(originals))
for _, itemMap := range originals {
q := tx.NewInsert().Table(tableName)
for key, value := range itemMap {
q = q.Value(key, value)
}
if pkName == "" {
if _, err := q.Exec(ctx); err != nil {
return err
}
originals = append(originals, itemMap)
insertedIDs = append(insertedIDs, returnedID)
insertedIDs = append(insertedIDs, nil)
continue
}
return nil
})
if err != nil {
var returnedID interface{}
if err := q.Returning(pkName).Scan(ctx, &returnedID); err != nil {
return err
}
insertedIDs = append(insertedIDs, returnedID)
}
return nil
})
if err != nil {
if single {
return nil, fmt.Errorf("create error: %w", err)
}
if _, ok := hookCtx.Data.([]interface{}); ok {
return nil, fmt.Errorf("batch create error: %w", err)
}
// Re-fetch each record after transaction commits; fall back to input on failure.
results := make([]interface{}, 0, len(insertedIDs))
return nil, err
}
// Transaction 2: re-fetch to capture DB-generated defaults/triggers, then AfterCreate.
results := make([]interface{}, 0, len(insertedIDs))
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
results = results[:0]
for i, pkVal := range insertedIDs {
if pkVal == nil {
results = append(results, originals[i])
continue
}
fetchedRecord := reflect.New(modelType).Interface()
if err := h.db.NewSelect().Model(fetchedRecord).
if err := tx.NewSelect().Model(fetchedRecord).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), pkVal).
ScanModel(ctx); err == nil {
results = append(results, mergeWithInput(fetchedRecord, originals[i]))
@@ -529,15 +533,23 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
results = append(results, originals[i])
}
}
hookCtx.Result = results
if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil {
return nil, fmt.Errorf("AfterCreate hook failed: %w", err)
if single {
hookCtx.Result = results[0]
} else {
hookCtx.Result = results
}
return results, nil
default:
return nil, fmt.Errorf("data must be an object or array of objects")
if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil {
return fmt.Errorf("AfterCreate hook failed: %w", err)
}
return nil
})
if err != nil {
return nil, err
}
if single {
return results[0], nil
}
return results, nil
}
// executeUpdate updates a record by ID.
@@ -573,8 +585,20 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
pkName := reflection.GetPrimaryKeyName(model)
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "update",
ID: id,
Data: updates,
Tx: h.db,
}
var updateResult interface{}
err = h.db.RunInTransaction(ctx, func(tx common.Database) error {
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
// Read existing record
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
@@ -601,17 +625,6 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
return fmt.Errorf("error unmarshaling existing record: %w", err)
}
hookCtx := &HookContext{
Context: ctx,
Handler: h,
Schema: schema,
Entity: entity,
Model: model,
Operation: "update",
ID: id,
Data: updates,
Tx: tx,
}
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
return err
}
@@ -649,22 +662,28 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
return nil, err
}
// Re-fetch the record after transaction commits to capture DB-generated changes.
// Transaction 2: re-fetch to capture DB-generated changes.
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
modelType = modelType.Elem()
}
fetchedRecord := reflect.New(modelType).Interface()
if err := h.db.NewSelect().Model(fetchedRecord).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id).
ScanModel(ctx); err == nil {
jsonData, marshalErr := json.Marshal(fetchedRecord)
if marshalErr == nil {
var fetchedMap map[string]interface{}
if json.Unmarshal(jsonData, &fetchedMap) == nil {
updateResult = fetchedMap
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
fetchedRecord := reflect.New(modelType).Interface()
if err := tx.NewSelect().Model(fetchedRecord).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id).
ScanModel(ctx); err == nil {
jsonData, marshalErr := json.Marshal(fetchedRecord)
if marshalErr == nil {
var fetchedMap map[string]interface{}
if json.Unmarshal(jsonData, &fetchedMap) == nil {
updateResult = fetchedMap
}
}
}
return nil
})
if err != nil {
return nil, err
}
return updateResult, nil
@@ -706,9 +725,6 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
return nil, err
}
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
return nil, err
}
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
@@ -717,7 +733,10 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
var recordToDelete interface{}
err = h.db.RunInTransaction(ctx, func(tx common.Database) error {
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
return err
}
record := reflect.New(modelType).Interface()
selectQuery := tx.NewSelect().Model(record).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
@@ -739,7 +758,6 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
}
recordToDelete = record
hookCtx.Tx = tx
hookCtx.Result = record
return h.hooks.Execute(AfterDelete, hookCtx)
})
@@ -873,3 +891,11 @@ func (h *Handler) applyPreloads(model interface{}, query common.SelectQuery, pre
}
return query, nil
}
// 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)
}
+9
View File
@@ -26,6 +26,12 @@ const (
BeforeDelete HookType = "before_delete"
AfterDelete HookType = "after_delete"
// 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
@@ -48,6 +54,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
type HookFunc func(*HookContext) error
+207
View File
@@ -0,0 +1,207 @@
package resolvemcp
import (
"context"
"database/sql"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
type txItem struct {
ID int `json:"id" bun:"id,pk"`
Name string `json:"name" bun:"name"`
}
func newTxHarness(t *testing.T) (*Handler, sqlmock.Sqlmock, context.Context) {
t.Helper()
db, mock, err := sqlmock.New()
if err != nil {
t.Fatal(err)
}
// One connection: any statement bypassing the open tx cannot get a
// connection and fails on the context timeout.
db.SetMaxOpenConns(1)
t.Cleanup(func() { _ = db.Close() })
h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry(), Config{})
if err := h.RegisterModel("public", "items", &txItem{}); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
t.Cleanup(cancel)
return h, mock, ctx
}
// trace records hook firing order and the Tx each hook saw.
type trace struct {
order []string
txs map[string][]common.Database
}
func traceHooks(h *Handler, types ...HookType) *trace {
tr := &trace{txs: map[string][]common.Database{}}
for _, ht := range types {
ht := ht
h.Hooks().Register(ht, func(c *HookContext) error {
tr.order = append(tr.order, string(ht))
tr.txs[string(ht)] = append(tr.txs[string(ht)], c.Tx)
return nil
})
}
return tr
}
func (tr *trace) assertOrder(t *testing.T, want ...string) {
t.Helper()
if len(tr.order) != len(want) {
t.Fatalf("hook order %v, want %v", tr.order, want)
}
for i := range want {
if tr.order[i] != want[i] {
t.Fatalf("hook order %v, want %v", tr.order, want)
}
}
}
func TestDeleteRunsHooksInOneTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, BeforeDelete, AfterDelete)
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
mock.ExpectExec(`DELETE`).WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
if _, err := h.executeDelete(ctx, "public", "items", "7"); err != nil {
t.Fatal(err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
tr.assertOrder(t, "on_tx_begin", "before_delete", "after_delete")
if tr.txs["before_delete"][0] != tr.txs["on_tx_begin"][0] || tr.txs["after_delete"][0] != tr.txs["on_tx_begin"][0] {
t.Fatal("OnTxBegin, BeforeDelete and AfterDelete must share one transaction")
}
}
func TestDeleteBeforeHookErrorRollsBack(t *testing.T) {
h, mock, ctx := newTxHarness(t)
h.Hooks().Register(BeforeDelete, func(*HookContext) error { return sql.ErrConnDone })
mock.ExpectBegin()
mock.ExpectRollback()
if _, err := h.executeDelete(ctx, "public", "items", "7"); err == nil {
t.Fatal("expected error")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestReadRunsInOneTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, BeforeRead, AfterRead)
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()
if _, _, err := h.executeRead(ctx, "public", "items", "7", common.RequestOptions{}); err != nil {
t.Fatal(err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
tr.assertOrder(t, "on_tx_begin", "before_read", "after_read")
for _, ht := range []string{"before_read", "after_read"} {
if tr.txs[ht][0] != tr.txs["on_tx_begin"][0] {
t.Fatalf("%s must run on the OnTxBegin transaction", ht)
}
}
}
func TestCreateSingleUsesTwoTransactions(t *testing.T) {
h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, BeforeCreate, AfterCreate)
mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
mock.ExpectCommit()
if _, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a"}); err != nil {
t.Fatal(err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
tr.assertOrder(t, "on_tx_begin", "before_create", "on_tx_begin", "after_create")
if tr.txs["before_create"][0] != tr.txs["on_tx_begin"][0] {
t.Fatal("BeforeCreate must run on the first transaction")
}
if tr.txs["after_create"][0] != tr.txs["on_tx_begin"][1] || tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] {
t.Fatal("AfterCreate must run on a second, distinct transaction")
}
}
func TestCreateBatchRefetchOnSecondTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, AfterCreate)
mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(2, "b"))
mock.ExpectCommit()
items := []interface{}{map[string]interface{}{"name": "a"}, map[string]interface{}{"name": "b"}}
if _, err := h.executeCreate(ctx, "public", "items", items); err != nil {
t.Fatal(err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
tr.assertOrder(t, "on_tx_begin", "on_tx_begin", "after_create")
}
func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(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()
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)
}
tr.assertOrder(t, "on_tx_begin", "before_update", "after_update", "on_tx_begin")
if tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] {
t.Fatal("re-fetch must run on a second transaction")
}
for _, ht := range []string{"before_update", "after_update"} {
if tr.txs[ht][0] != tr.txs["on_tx_begin"][0] {
t.Fatalf("%s must run on the first transaction", ht)
}
}
}