mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-08 14:26:28 +00:00
feat(tx): run resolvemcp operations in per-operation transactions with OnTxBegin
This commit is contained in:
@@ -76,7 +76,7 @@
|
|||||||
| 2 | DONE | `OnTxBegin` hook type + `runInTx` helper | `common/txhook.go`, `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec; other specs in P4-6 |
|
| 2 | DONE | `OnTxBegin` hook type + `runInTx` helper | `common/txhook.go`, `*/hooks.go`, `*/handler.go` | resolvespec + restheadspec; other specs in P4-6 |
|
||||||
| 3 | DONE | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above |
|
| 3 | DONE | Insert/update post-commit hooks + re-fetch in second short `runInTx` (select only) | restheadspec `:1005, 1467, 1667-1674`; resolvespec `:1297, 1449, 1602` | per decision above |
|
||||||
| 4 | DONE | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias |
|
| 4 | DONE | websocketspec + mqttspec: wrap read/create/update/delete in `runInTx` | `websocketspec/handler.go`, `mqttspec/handler.go` | mqttspec aliases websocketspec hooks; confirm `OnTxBegin` alias |
|
||||||
| 5 | TODO | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx |
|
| 5 | DONE | resolvemcp: read + single create in tx | `resolvemcp/handler.go:253, 445` | verify batch/update/delete hooks run inside tx |
|
||||||
| 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | |
|
| 6 | TODO | funcspec: `OnTxBegin` (or once-per-tx `BeforeOp`), `BeforeResponse` via `runInTx` | `funcspec/function_api.go:337, 640` | |
|
||||||
| 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | |
|
| 7 | TODO | Security hooks: register RLS stamping on `OnTxBegin`; document | `pkg/security/*`, README | |
|
||||||
|
|
||||||
@@ -91,7 +91,8 @@
|
|||||||
- OPEN: restheadspec `AfterRead` still runs post-commit with `Tx = h.db` (`:1004`); decision says read has no second tx. Needs a call: run it inside the read tx, or in a short second tx.
|
- OPEN: restheadspec `AfterRead` still runs post-commit with `Tx = h.db` (`:1004`); decision says read has no second tx. Needs a call: run it inside the read tx, or in a short second tx.
|
||||||
- DONE P4: websocketspec + mqttspec. `OnTxBegin` (mqttspec re-exports the websocketspec constant), `HookContext.SetTx`, per-handler `runInTx`/`sendTxError`. Per message: read = 1 tx (Before/After hooks + queries); delete = 1 tx (Before, delete, After); create/update = tx 1 (Before + write) then tx 2 (re-fetch + `BeforeScan` + After). `create()`/`update()` no longer re-fetch; `read*`/`create`/`update`/`delete` use `hookCtx.Tx`. websocketspec `FetchRowNumber` keeps its public signature and delegates to a new tx-aware `fetchRowNumber`. A failure in begin/`OnTxBegin`/commit answers `transaction_error` with no detail. Tests: `pkg/websocketspec/tx_test.go` (sqlmock), `pkg/mqttspec/tx_test.go` (sqlite); mqttspec `update` tests now pass `Tx`.
|
- DONE P4: websocketspec + mqttspec. `OnTxBegin` (mqttspec re-exports the websocketspec constant), `HookContext.SetTx`, per-handler `runInTx`/`sendTxError`. Per message: read = 1 tx (Before/After hooks + queries); delete = 1 tx (Before, delete, After); create/update = tx 1 (Before + write) then tx 2 (re-fetch + `BeforeScan` + After). `create()`/`update()` no longer re-fetch; `read*`/`create`/`update`/`delete` use `hookCtx.Tx`. websocketspec `FetchRowNumber` keeps its public signature and delegates to a new tx-aware `fetchRowNumber`. A failure in begin/`OnTxBegin`/commit answers `transaction_error` with no detail. Tests: `pkg/websocketspec/tx_test.go` (sqlmock), `pkg/mqttspec/tx_test.go` (sqlite); mqttspec `update` tests now pass `Tx`.
|
||||||
- DECIDED in P4 (follow `AfterRead` question above): websocketspec/mqttspec run `AfterRead` inside the read tx (keeps "read has no second tx").
|
- DECIDED in P4 (follow `AfterRead` question above): websocketspec/mqttspec run `AfterRead` inside the read tx (keeps "read has no second tx").
|
||||||
- NEXT: P5.
|
- DONE P5: resolvemcp. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx`. Read = 1 tx (`BeforeRead`, count, scan, `AfterRead`; `readInTx`). Delete = 1 tx (`BeforeDelete` moved inside, after `OnTxBegin`). Create (single and batch, unified) = tx 1 (`BeforeCreate` + inserts) then tx 2 (re-fetch + `AfterCreate`); the old single-record pool insert/re-fetch is gone. Update = tx 1 (select, `BeforeUpdate`, update, `AfterUpdate`) then tx 2 (re-fetch). `BeforeHandle` still runs before any tx with `Tx = h.db`. Tests: `pkg/resolvemcp/tx_test.go` (sqlmock).
|
||||||
|
- NEXT: P6.
|
||||||
|
|
||||||
## Tests
|
## Tests
|
||||||
- Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests.
|
- Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests.
|
||||||
|
|||||||
+131
-105
@@ -247,10 +247,27 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
|
|||||||
return nil, nil, err
|
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))
|
sliceType := reflect.SliceOf(reflect.PointerTo(modelType))
|
||||||
modelPtr := reflect.New(sliceType).Interface()
|
modelPtr := reflect.New(sliceType).Interface()
|
||||||
|
|
||||||
query := h.db.NewSelect().Model(modelPtr)
|
query := hookCtx.Tx.NewSelect().Model(modelPtr)
|
||||||
|
|
||||||
tempInstance := reflect.New(modelType).Interface()
|
tempInstance := reflect.New(modelType).Interface()
|
||||||
if provider, ok := tempInstance.(common.TableNameProvider); !ok || provider.TableName() == "" {
|
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 {
|
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
|
||||||
return nil, err
|
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)
|
pkName := reflection.GetPrimaryKeyName(model)
|
||||||
|
modelType := reflect.TypeOf(model)
|
||||||
|
if modelType.Kind() == reflect.Pointer {
|
||||||
|
modelType = modelType.Elem()
|
||||||
|
}
|
||||||
|
|
||||||
switch v := data.(type) {
|
// Transaction 1: BeforeCreate + inserts.
|
||||||
case map[string]interface{}:
|
var (
|
||||||
query := h.db.NewInsert().Table(tableName)
|
single bool
|
||||||
for key, value := range v {
|
originals []map[string]interface{}
|
||||||
query = query.Value(key, value)
|
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 != "" {
|
// Use potentially modified data
|
||||||
var insertedID interface{}
|
switch v := hookCtx.Data.(type) {
|
||||||
if err := query.Returning(pkName).Scan(ctx, &insertedID); err != nil {
|
case map[string]interface{}:
|
||||||
return nil, fmt.Errorf("create error: %w", err)
|
single = true
|
||||||
}
|
originals = []map[string]interface{}{v}
|
||||||
// Re-fetch after insert to capture DB-generated defaults/triggers.
|
case []interface{}:
|
||||||
modelType := reflect.TypeOf(model)
|
originals = make([]map[string]interface{}, 0, len(v))
|
||||||
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 {
|
|
||||||
for _, item := range v {
|
for _, item := range v {
|
||||||
itemMap, ok := item.(map[string]interface{})
|
itemMap, ok := item.(map[string]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return fmt.Errorf("each item must be an object")
|
return fmt.Errorf("each item must be an object")
|
||||||
}
|
}
|
||||||
q := tx.NewInsert().Table(tableName)
|
originals = append(originals, itemMap)
|
||||||
for key, value := range itemMap {
|
}
|
||||||
q = q.Value(key, value)
|
default:
|
||||||
}
|
return fmt.Errorf("data must be an object or array of objects")
|
||||||
if pkName == "" {
|
}
|
||||||
if _, err := q.Exec(ctx); err != nil {
|
|
||||||
return err
|
insertedIDs = make([]interface{}, 0, len(originals))
|
||||||
}
|
for _, itemMap := range originals {
|
||||||
originals = append(originals, itemMap)
|
q := tx.NewInsert().Table(tableName)
|
||||||
insertedIDs = append(insertedIDs, nil)
|
for key, value := range itemMap {
|
||||||
continue
|
q = q.Value(key, value)
|
||||||
}
|
}
|
||||||
var returnedID interface{}
|
if pkName == "" {
|
||||||
if err := q.Returning(pkName).Scan(ctx, &returnedID); err != nil {
|
if _, err := q.Exec(ctx); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
originals = append(originals, itemMap)
|
insertedIDs = append(insertedIDs, nil)
|
||||||
insertedIDs = append(insertedIDs, returnedID)
|
continue
|
||||||
}
|
}
|
||||||
return nil
|
var returnedID interface{}
|
||||||
})
|
if err := q.Returning(pkName).Scan(ctx, &returnedID); err != nil {
|
||||||
if 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)
|
return nil, fmt.Errorf("batch create error: %w", err)
|
||||||
}
|
}
|
||||||
// Re-fetch each record after transaction commits; fall back to input on failure.
|
return nil, err
|
||||||
results := make([]interface{}, 0, len(insertedIDs))
|
}
|
||||||
|
|
||||||
|
// 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 {
|
for i, pkVal := range insertedIDs {
|
||||||
if pkVal == nil {
|
if pkVal == nil {
|
||||||
results = append(results, originals[i])
|
results = append(results, originals[i])
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
fetchedRecord := reflect.New(modelType).Interface()
|
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).
|
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), pkVal).
|
||||||
ScanModel(ctx); err == nil {
|
ScanModel(ctx); err == nil {
|
||||||
results = append(results, mergeWithInput(fetchedRecord, originals[i]))
|
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])
|
results = append(results, originals[i])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
hookCtx.Result = results
|
if single {
|
||||||
if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil {
|
hookCtx.Result = results[0]
|
||||||
return nil, fmt.Errorf("AfterCreate hook failed: %w", err)
|
} else {
|
||||||
|
hookCtx.Result = results
|
||||||
}
|
}
|
||||||
return results, nil
|
if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil {
|
||||||
|
return fmt.Errorf("AfterCreate hook failed: %w", err)
|
||||||
default:
|
}
|
||||||
return nil, fmt.Errorf("data must be an object or array of objects")
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if single {
|
||||||
|
return results[0], nil
|
||||||
|
}
|
||||||
|
return results, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// executeUpdate updates a record by ID.
|
// 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)
|
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{}
|
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
|
// Read existing record
|
||||||
modelType := reflect.TypeOf(model)
|
modelType := reflect.TypeOf(model)
|
||||||
if modelType.Kind() == reflect.Pointer {
|
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)
|
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 {
|
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -649,22 +662,28 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
|
|||||||
return nil, err
|
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)
|
modelType := reflect.TypeOf(model)
|
||||||
if modelType.Kind() == reflect.Pointer {
|
if modelType.Kind() == reflect.Pointer {
|
||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
fetchedRecord := reflect.New(modelType).Interface()
|
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||||
if err := h.db.NewSelect().Model(fetchedRecord).
|
fetchedRecord := reflect.New(modelType).Interface()
|
||||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id).
|
if err := tx.NewSelect().Model(fetchedRecord).
|
||||||
ScanModel(ctx); err == nil {
|
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id).
|
||||||
jsonData, marshalErr := json.Marshal(fetchedRecord)
|
ScanModel(ctx); err == nil {
|
||||||
if marshalErr == nil {
|
jsonData, marshalErr := json.Marshal(fetchedRecord)
|
||||||
var fetchedMap map[string]interface{}
|
if marshalErr == nil {
|
||||||
if json.Unmarshal(jsonData, &fetchedMap) == nil {
|
var fetchedMap map[string]interface{}
|
||||||
updateResult = fetchedMap
|
if json.Unmarshal(jsonData, &fetchedMap) == nil {
|
||||||
|
updateResult = fetchedMap
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
return updateResult, nil
|
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 {
|
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := h.hooks.Execute(BeforeDelete, hookCtx); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
modelType := reflect.TypeOf(model)
|
modelType := reflect.TypeOf(model)
|
||||||
if modelType.Kind() == reflect.Pointer {
|
if modelType.Kind() == reflect.Pointer {
|
||||||
@@ -717,7 +733,10 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
|
|||||||
|
|
||||||
var recordToDelete interface{}
|
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()
|
record := reflect.New(modelType).Interface()
|
||||||
selectQuery := tx.NewSelect().Model(record).
|
selectQuery := tx.NewSelect().Model(record).
|
||||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
|
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
|
recordToDelete = record
|
||||||
hookCtx.Tx = tx
|
|
||||||
hookCtx.Result = record
|
hookCtx.Result = record
|
||||||
return h.hooks.Execute(AfterDelete, hookCtx)
|
return h.hooks.Execute(AfterDelete, hookCtx)
|
||||||
})
|
})
|
||||||
@@ -873,3 +891,11 @@ func (h *Handler) applyPreloads(model interface{}, query common.SelectQuery, pre
|
|||||||
}
|
}
|
||||||
return query, nil
|
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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -26,6 +26,12 @@ const (
|
|||||||
|
|
||||||
BeforeDelete HookType = "before_delete"
|
BeforeDelete HookType = "before_delete"
|
||||||
AfterDelete HookType = "after_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
|
// HookContext contains all the data available to a hook
|
||||||
@@ -48,6 +54,9 @@ type HookContext struct {
|
|||||||
Tx common.Database
|
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
|
// HookFunc is the signature for hook functions
|
||||||
type HookFunc func(*HookContext) error
|
type HookFunc func(*HookContext) error
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user