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)
}