mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 19:41:57 +00:00
feat(tx): run resolvemcp operations in per-operation transactions with OnTxBegin
This commit is contained in:
+131
-105
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user