mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-03 03:51:59 +00:00
feat(tx): run funcspec OnTxBegin and BeforeResponse in transactions
This commit is contained in:
@@ -77,7 +77,7 @@
|
|||||||
| 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 | DONE | 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 | DONE | 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 | |
|
||||||
|
|
||||||
## Progress
|
## Progress
|
||||||
@@ -92,7 +92,8 @@
|
|||||||
- 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").
|
||||||
- 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).
|
- 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.
|
- DONE P6: funcspec. `OnTxBegin`, `HookContext.SetTx`, `Handler.runInTx` for `SqlQuery` and `SqlQueryList`. `BeforeResponse` now runs in a second short tx (`Tx` is no longer the pool). `BeforeOp` is unchanged (still per statement). A begin/`OnTxBegin`/commit failure answers 500 `transaction_error` / "Transaction failed" (before, it returned with no response); body failures still answer via `sendError`. Tests: `pkg/funcspec/tx_test.go`.
|
||||||
|
- NEXT: P7.
|
||||||
|
|
||||||
## 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.
|
||||||
|
|||||||
+207
-181
@@ -192,131 +192,140 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
hookCtx.InputVars = inputvars
|
hookCtx.InputVars = inputvars
|
||||||
|
|
||||||
// Execute query within transaction
|
// Execute query within transaction
|
||||||
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
|
// bodyRan/bodyFailed tell a begin/OnTxBegin/commit failure (no response sent
|
||||||
// Set transaction in hook context for hooks to use
|
// yet) from a body failure (sendError already answered).
|
||||||
hookCtx.Tx = tx
|
var bodyRan, bodyFailed bool
|
||||||
|
err := h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||||
|
bodyRan = true
|
||||||
|
berr := func() error {
|
||||||
|
|
||||||
// Execute BeforeQueryList hook (inside transaction)
|
// Execute BeforeQueryList hook (inside transaction)
|
||||||
if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil {
|
||||||
logger.Error("BeforeQueryList hook failed: %v", err)
|
logger.Error("BeforeQueryList hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check if hook aborted the operation
|
|
||||||
if hookCtx.Abort {
|
|
||||||
if hookCtx.AbortCode == 0 {
|
|
||||||
hookCtx.AbortCode = http.StatusBadRequest
|
|
||||||
}
|
|
||||||
sendError(w, hookCtx.AbortCode, "operation_aborted", hookCtx.AbortMessage, nil)
|
|
||||||
return fmt.Errorf("operation aborted: %s", hookCtx.AbortMessage)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Use potentially modified SQL query from hook
|
|
||||||
sqlquery = hookCtx.SQLQuery
|
|
||||||
sqlqueryCnt := sqlquery
|
|
||||||
|
|
||||||
// Parse sorting and pagination parameters
|
|
||||||
sortcols, limit, offset := h.parsePaginationParams(r)
|
|
||||||
|
|
||||||
// Override with parsed parameters if available
|
|
||||||
if reqParams.SortColumns != "" {
|
|
||||||
sortcols = reqParams.SortColumns
|
|
||||||
}
|
|
||||||
if reqParams.Limit > 0 {
|
|
||||||
limit = reqParams.Limit
|
|
||||||
}
|
|
||||||
if reqParams.Offset > 0 {
|
|
||||||
offset = reqParams.Offset
|
|
||||||
}
|
|
||||||
|
|
||||||
hookCtx.SortColumns = sortcols
|
|
||||||
hookCtx.Limit = limit
|
|
||||||
hookCtx.Offset = offset
|
|
||||||
fromPos := strings.Index(strings.ToLower(sqlquery), "from ")
|
|
||||||
orderbyPos := strings.Index(strings.ToLower(sqlquery), "order by")
|
|
||||||
|
|
||||||
if len(sortcols) > 0 && (orderbyPos < 0 || (orderbyPos > 0 && orderbyPos < fromPos)) {
|
|
||||||
sqlquery = fmt.Sprintf("%s \nORDER BY %s", sqlquery, ValidSQL(sortcols, "select"))
|
|
||||||
}
|
|
||||||
|
|
||||||
if !options.NoCount {
|
|
||||||
if limit > 0 && offset > 0 {
|
|
||||||
sqlquery = fmt.Sprintf("%s \nLIMIT %d OFFSET %d", sqlquery, limit, offset)
|
|
||||||
} else if limit > 0 {
|
|
||||||
sqlquery = fmt.Sprintf("%s \nLIMIT %d", sqlquery, limit)
|
|
||||||
} else {
|
|
||||||
sqlquery = fmt.Sprintf("%s \nLIMIT %d", sqlquery, 20000)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Get total count
|
|
||||||
countQuery := fmt.Sprintf("SELECT COUNT(1) FROM (%s) cnts", sqlqueryCnt)
|
|
||||||
var countResult struct{ Count int64 }
|
|
||||||
if err := tx.Query(ctx, &countResult, countQuery); err != nil {
|
|
||||||
sendError(w, http.StatusBadRequest, "count_failed", "Failed to retrieve record count", err)
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
total = countResult.Count
|
|
||||||
}
|
|
||||||
|
|
||||||
// Execute BeforeSQLExec hook
|
// Check if hook aborted the operation
|
||||||
hookCtx.SQLQuery = sqlquery
|
if hookCtx.Abort {
|
||||||
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
if hookCtx.AbortCode == 0 {
|
||||||
logger.Error("BeforeSQLExec hook failed: %v", err)
|
hookCtx.AbortCode = http.StatusBadRequest
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
}
|
||||||
return err
|
sendError(w, hookCtx.AbortCode, "operation_aborted", hookCtx.AbortMessage, nil)
|
||||||
}
|
return fmt.Errorf("operation aborted: %s", hookCtx.AbortMessage)
|
||||||
// Use potentially modified SQL query from hook
|
}
|
||||||
sqlquery = hookCtx.SQLQuery
|
|
||||||
|
|
||||||
// Execute main query
|
// Use potentially modified SQL query from hook
|
||||||
rows := make([]map[string]interface{}, 0)
|
sqlquery = hookCtx.SQLQuery
|
||||||
if err := tx.Query(ctx, &rows, sqlquery); err != nil {
|
sqlqueryCnt := sqlquery
|
||||||
sendError(w, http.StatusBadRequest, "query_failed", "Failed to retrieve records", err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Normalize PostgreSQL types for proper JSON marshaling
|
// Parse sorting and pagination parameters
|
||||||
dbobjlist = normalizePostgresTypesList(rows)
|
sortcols, limit, offset := h.parsePaginationParams(r)
|
||||||
|
|
||||||
if options.NoCount {
|
// Override with parsed parameters if available
|
||||||
total = int64(len(dbobjlist))
|
if reqParams.SortColumns != "" {
|
||||||
}
|
sortcols = reqParams.SortColumns
|
||||||
|
}
|
||||||
|
if reqParams.Limit > 0 {
|
||||||
|
limit = reqParams.Limit
|
||||||
|
}
|
||||||
|
if reqParams.Offset > 0 {
|
||||||
|
offset = reqParams.Offset
|
||||||
|
}
|
||||||
|
|
||||||
// Execute AfterSQLExec hook
|
hookCtx.SortColumns = sortcols
|
||||||
hookCtx.Result = dbobjlist
|
hookCtx.Limit = limit
|
||||||
hookCtx.Total = total
|
hookCtx.Offset = offset
|
||||||
if err := h.hooks.Execute(AfterSQLExec, hookCtx); err != nil {
|
fromPos := strings.Index(strings.ToLower(sqlquery), "from ")
|
||||||
logger.Error("AfterSQLExec hook failed: %v", err)
|
orderbyPos := strings.Index(strings.ToLower(sqlquery), "order by")
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
// Use potentially modified result from hook
|
|
||||||
if modifiedResult, ok := hookCtx.Result.([]map[string]interface{}); ok {
|
|
||||||
dbobjlist = modifiedResult
|
|
||||||
}
|
|
||||||
total = hookCtx.Total
|
|
||||||
|
|
||||||
// Execute AfterQueryList hook (inside transaction)
|
if len(sortcols) > 0 && (orderbyPos < 0 || (orderbyPos > 0 && orderbyPos < fromPos)) {
|
||||||
hookCtx.Result = dbobjlist
|
sqlquery = fmt.Sprintf("%s \nORDER BY %s", sqlquery, ValidSQL(sortcols, "select"))
|
||||||
hookCtx.Total = total
|
}
|
||||||
hookCtx.Error = nil
|
|
||||||
if err := h.hooks.Execute(AfterQueryList, hookCtx); err != nil {
|
|
||||||
logger.Error("AfterQueryList hook failed: %v", err)
|
|
||||||
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
// Use potentially modified result from hook
|
|
||||||
if modifiedResult, ok := hookCtx.Result.([]map[string]interface{}); ok {
|
|
||||||
dbobjlist = modifiedResult
|
|
||||||
}
|
|
||||||
total = hookCtx.Total
|
|
||||||
|
|
||||||
return nil
|
if !options.NoCount {
|
||||||
|
if limit > 0 && offset > 0 {
|
||||||
|
sqlquery = fmt.Sprintf("%s \nLIMIT %d OFFSET %d", sqlquery, limit, offset)
|
||||||
|
} else if limit > 0 {
|
||||||
|
sqlquery = fmt.Sprintf("%s \nLIMIT %d", sqlquery, limit)
|
||||||
|
} else {
|
||||||
|
sqlquery = fmt.Sprintf("%s \nLIMIT %d", sqlquery, 20000)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get total count
|
||||||
|
countQuery := fmt.Sprintf("SELECT COUNT(1) FROM (%s) cnts", sqlqueryCnt)
|
||||||
|
var countResult struct{ Count int64 }
|
||||||
|
if err := tx.Query(ctx, &countResult, countQuery); err != nil {
|
||||||
|
sendError(w, http.StatusBadRequest, "count_failed", "Failed to retrieve record count", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
total = countResult.Count
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute BeforeSQLExec hook
|
||||||
|
hookCtx.SQLQuery = sqlquery
|
||||||
|
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
||||||
|
logger.Error("BeforeSQLExec hook failed: %v", err)
|
||||||
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// Use potentially modified SQL query from hook
|
||||||
|
sqlquery = hookCtx.SQLQuery
|
||||||
|
|
||||||
|
// Execute main query
|
||||||
|
rows := make([]map[string]interface{}, 0)
|
||||||
|
if err := tx.Query(ctx, &rows, sqlquery); err != nil {
|
||||||
|
sendError(w, http.StatusBadRequest, "query_failed", "Failed to retrieve records", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Normalize PostgreSQL types for proper JSON marshaling
|
||||||
|
dbobjlist = normalizePostgresTypesList(rows)
|
||||||
|
|
||||||
|
if options.NoCount {
|
||||||
|
total = int64(len(dbobjlist))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute AfterSQLExec hook
|
||||||
|
hookCtx.Result = dbobjlist
|
||||||
|
hookCtx.Total = total
|
||||||
|
if err := h.hooks.Execute(AfterSQLExec, hookCtx); err != nil {
|
||||||
|
logger.Error("AfterSQLExec hook failed: %v", err)
|
||||||
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// Use potentially modified result from hook
|
||||||
|
if modifiedResult, ok := hookCtx.Result.([]map[string]interface{}); ok {
|
||||||
|
dbobjlist = modifiedResult
|
||||||
|
}
|
||||||
|
total = hookCtx.Total
|
||||||
|
|
||||||
|
// Execute AfterQueryList hook (inside transaction)
|
||||||
|
hookCtx.Result = dbobjlist
|
||||||
|
hookCtx.Total = total
|
||||||
|
hookCtx.Error = nil
|
||||||
|
if err := h.hooks.Execute(AfterQueryList, hookCtx); err != nil {
|
||||||
|
logger.Error("AfterQueryList hook failed: %v", err)
|
||||||
|
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// Use potentially modified result from hook
|
||||||
|
if modifiedResult, ok := hookCtx.Result.([]map[string]interface{}); ok {
|
||||||
|
dbobjlist = modifiedResult
|
||||||
|
}
|
||||||
|
total = hookCtx.Total
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}()
|
||||||
|
bodyFailed = berr != nil
|
||||||
|
return berr
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("Transaction failed: %v", err)
|
logger.Error("Transaction failed: %v", err)
|
||||||
|
if !bodyRan || !bodyFailed {
|
||||||
|
sendError(w, http.StatusInternalServerError, "transaction_error", "Transaction failed", nil)
|
||||||
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -331,13 +340,13 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
|
|||||||
w.Header().Set("Content-Range", fmt.Sprintf("items %d-%d/%d", respOffset, respOffset+len(dbobjlist), total))
|
w.Header().Set("Content-Range", fmt.Sprintf("items %d-%d/%d", respOffset, respOffset+len(dbobjlist), total))
|
||||||
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
|
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
|
||||||
|
|
||||||
// Execute BeforeResponse hook. The transaction has already committed by
|
// Execute BeforeResponse hook in a second short transaction: the main one
|
||||||
// this point, so hooks must use the pooled connection rather than the
|
// has already committed, and hooks must never get the pooled connection.
|
||||||
// now-dead tx.
|
|
||||||
hookCtx.Tx = h.db
|
|
||||||
hookCtx.Result = dbobjlist
|
hookCtx.Result = dbobjlist
|
||||||
hookCtx.Total = total
|
hookCtx.Total = total
|
||||||
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
if err := h.runInTx(ctx, hookCtx, func(common.Database) error {
|
||||||
|
return h.hooks.Execute(BeforeResponse, hookCtx)
|
||||||
|
}); err != nil {
|
||||||
logger.Error("BeforeResponse hook failed: %v", err)
|
logger.Error("BeforeResponse hook failed: %v", err)
|
||||||
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
||||||
return
|
return
|
||||||
@@ -558,88 +567,97 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
|
|||||||
hookCtx.InputVars = inputvars
|
hookCtx.InputVars = inputvars
|
||||||
|
|
||||||
// Execute query within transaction
|
// Execute query within transaction
|
||||||
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
|
// bodyRan/bodyFailed tell a begin/OnTxBegin/commit failure (no response sent
|
||||||
// Set transaction in hook context for hooks to use
|
// yet) from a body failure (sendError already answered).
|
||||||
hookCtx.Tx = tx
|
var bodyRan, bodyFailed bool
|
||||||
|
err := h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||||
|
bodyRan = true
|
||||||
|
berr := func() error {
|
||||||
|
|
||||||
// Execute BeforeQuery hook (inside transaction)
|
// Execute BeforeQuery hook (inside transaction)
|
||||||
if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil {
|
if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil {
|
||||||
logger.Error("BeforeQuery hook failed: %v", err)
|
logger.Error("BeforeQuery hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
}
|
|
||||||
|
|
||||||
// Check if hook aborted the operation
|
|
||||||
if hookCtx.Abort {
|
|
||||||
if hookCtx.AbortCode == 0 {
|
|
||||||
hookCtx.AbortCode = http.StatusBadRequest
|
|
||||||
}
|
}
|
||||||
sendError(w, hookCtx.AbortCode, "operation_aborted", hookCtx.AbortMessage, nil)
|
|
||||||
return fmt.Errorf("operation aborted: %s", hookCtx.AbortMessage)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Use potentially modified SQL query from hook
|
// Check if hook aborted the operation
|
||||||
sqlquery = hookCtx.SQLQuery
|
if hookCtx.Abort {
|
||||||
|
if hookCtx.AbortCode == 0 {
|
||||||
|
hookCtx.AbortCode = http.StatusBadRequest
|
||||||
|
}
|
||||||
|
sendError(w, hookCtx.AbortCode, "operation_aborted", hookCtx.AbortMessage, nil)
|
||||||
|
return fmt.Errorf("operation aborted: %s", hookCtx.AbortMessage)
|
||||||
|
}
|
||||||
|
|
||||||
// Execute BeforeSQLExec hook
|
// Use potentially modified SQL query from hook
|
||||||
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
sqlquery = hookCtx.SQLQuery
|
||||||
logger.Error("BeforeSQLExec hook failed: %v", err)
|
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
// Use potentially modified SQL query from hook
|
|
||||||
sqlquery = hookCtx.SQLQuery
|
|
||||||
|
|
||||||
// Execute main query
|
// Execute BeforeSQLExec hook
|
||||||
rows := make([]map[string]interface{}, 0)
|
if err := h.hooks.ExecuteBeforeOp(BeforeSQLExec, hookCtx); err != nil {
|
||||||
if err := tx.Query(ctx, &rows, sqlquery); err != nil {
|
logger.Error("BeforeSQLExec hook failed: %v", err)
|
||||||
sendError(w, http.StatusBadRequest, "query_failed", "Failed to retrieve records", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
// Use potentially modified SQL query from hook
|
||||||
|
sqlquery = hookCtx.SQLQuery
|
||||||
|
|
||||||
if len(rows) > 0 {
|
// Execute main query
|
||||||
dbobj = normalizePostgresTypes(rows[0])
|
rows := make([]map[string]interface{}, 0)
|
||||||
}
|
if err := tx.Query(ctx, &rows, sqlquery); err != nil {
|
||||||
|
sendError(w, http.StatusBadRequest, "query_failed", "Failed to retrieve records", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
// Execute AfterSQLExec hook
|
if len(rows) > 0 {
|
||||||
hookCtx.Result = dbobj
|
dbobj = normalizePostgresTypes(rows[0])
|
||||||
if err := h.hooks.Execute(AfterSQLExec, hookCtx); err != nil {
|
}
|
||||||
logger.Error("AfterSQLExec hook failed: %v", err)
|
|
||||||
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
// Use potentially modified result from hook
|
|
||||||
if modifiedResult, ok := hookCtx.Result.(map[string]interface{}); ok {
|
|
||||||
dbobj = modifiedResult
|
|
||||||
}
|
|
||||||
|
|
||||||
// Execute AfterQuery hook (inside transaction)
|
// Execute AfterSQLExec hook
|
||||||
hookCtx.Result = dbobj
|
hookCtx.Result = dbobj
|
||||||
hookCtx.Error = nil
|
if err := h.hooks.Execute(AfterSQLExec, hookCtx); err != nil {
|
||||||
if err := h.hooks.Execute(AfterQuery, hookCtx); err != nil {
|
logger.Error("AfterSQLExec hook failed: %v", err)
|
||||||
logger.Error("AfterQuery hook failed: %v", err)
|
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", err)
|
||||||
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
return err
|
||||||
return err
|
}
|
||||||
}
|
// Use potentially modified result from hook
|
||||||
// Use potentially modified result from hook
|
if modifiedResult, ok := hookCtx.Result.(map[string]interface{}); ok {
|
||||||
if modifiedResult, ok := hookCtx.Result.(map[string]interface{}); ok {
|
dbobj = modifiedResult
|
||||||
dbobj = modifiedResult
|
}
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
// Execute AfterQuery hook (inside transaction)
|
||||||
|
hookCtx.Result = dbobj
|
||||||
|
hookCtx.Error = nil
|
||||||
|
if err := h.hooks.Execute(AfterQuery, hookCtx); err != nil {
|
||||||
|
logger.Error("AfterQuery hook failed: %v", err)
|
||||||
|
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// Use potentially modified result from hook
|
||||||
|
if modifiedResult, ok := hookCtx.Result.(map[string]interface{}); ok {
|
||||||
|
dbobj = modifiedResult
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}()
|
||||||
|
bodyFailed = berr != nil
|
||||||
|
return berr
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.Error("Transaction failed: %v", err)
|
logger.Error("Transaction failed: %v", err)
|
||||||
|
if !bodyRan || !bodyFailed {
|
||||||
|
sendError(w, http.StatusInternalServerError, "transaction_error", "Transaction failed", nil)
|
||||||
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute BeforeResponse hook. The transaction has already committed by
|
// Execute BeforeResponse hook in a second short transaction: the main one
|
||||||
// this point, so hooks must use the pooled connection rather than the
|
// has already committed, and hooks must never get the pooled connection.
|
||||||
// now-dead tx.
|
|
||||||
hookCtx.Tx = h.db
|
|
||||||
hookCtx.Result = dbobj
|
hookCtx.Result = dbobj
|
||||||
if err := h.hooks.Execute(BeforeResponse, hookCtx); err != nil {
|
if err := h.runInTx(ctx, hookCtx, func(common.Database) error {
|
||||||
|
return h.hooks.Execute(BeforeResponse, hookCtx)
|
||||||
|
}); err != nil {
|
||||||
logger.Error("BeforeResponse hook failed: %v", err)
|
logger.Error("BeforeResponse hook failed: %v", err)
|
||||||
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
|
||||||
return
|
return
|
||||||
@@ -1249,3 +1267,11 @@ func normalizePostgresValue(value interface{}) interface{} {
|
|||||||
return v
|
return v
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -32,6 +32,12 @@ const (
|
|||||||
// BeforeOp fires immediately before every SQL operation (query, query list, SQL exec).
|
// BeforeOp fires immediately before every SQL operation (query, query list, SQL exec).
|
||||||
// It fires at each individual SQL-operation hook point, so it runs once per statement executed.
|
// It fires at each individual SQL-operation hook point, so it runs once per statement executed.
|
||||||
BeforeOp HookType = "before_op"
|
BeforeOp HookType = "before_op"
|
||||||
|
|
||||||
|
// OnTxBegin fires once, first, inside every transaction the handler opens
|
||||||
|
// (the main transaction and the short one that runs BeforeResponse).
|
||||||
|
// 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
|
||||||
@@ -75,6 +81,9 @@ type HookContext struct {
|
|||||||
AbortCode int // HTTP status code if aborted
|
AbortCode int // HTTP status code if aborted
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
// It receives a HookContext and can modify it or return an error
|
// It receives a HookContext and can modify it or return an error
|
||||||
// If an error is returned, the operation will be aborted
|
// If an error is returned, the operation will be aborted
|
||||||
|
|||||||
@@ -0,0 +1,102 @@
|
|||||||
|
package funcspec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
)
|
||||||
|
|
||||||
|
// txFactory returns a pool whose every transaction is a distinct MockDatabase.
|
||||||
|
func txFactory(queries *int) *MockDatabase {
|
||||||
|
return &MockDatabase{
|
||||||
|
RunInTransactionFunc: func(ctx context.Context, fn func(common.Database) error) error {
|
||||||
|
return fn(&MockDatabase{
|
||||||
|
QueryFunc: func(ctx context.Context, dest interface{}, query string, args ...interface{}) error {
|
||||||
|
*queries++
|
||||||
|
if rows, ok := dest.(*[]map[string]interface{}); ok {
|
||||||
|
*rows = []map[string]interface{}{{"id": float64(1)}}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOnTxBeginFirstAndBeforeResponseOnSecondTx(t *testing.T) {
|
||||||
|
var queries int
|
||||||
|
h := NewHandler(txFactory(&queries))
|
||||||
|
var order []string
|
||||||
|
txs := map[HookType][]common.Database{}
|
||||||
|
for _, ht := range []HookType{OnTxBegin, BeforeQuery, AfterQuery, BeforeResponse} {
|
||||||
|
ht := ht
|
||||||
|
h.Hooks().Register(ht, func(c *HookContext) error {
|
||||||
|
order = append(order, string(ht))
|
||||||
|
txs[ht] = append(txs[ht], c.Tx)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{})(w, createTestRequest("GET", "/t", nil, nil, nil))
|
||||||
|
|
||||||
|
if w.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status %d body %s", w.Code, w.Body)
|
||||||
|
}
|
||||||
|
if got := strings.Join(order, ","); got != "on_tx_begin,before_query,after_query,on_tx_begin,before_response" {
|
||||||
|
t.Fatalf("hook order %s", got)
|
||||||
|
}
|
||||||
|
if txs[BeforeQuery][0] != txs[OnTxBegin][0] || txs[AfterQuery][0] != txs[OnTxBegin][0] {
|
||||||
|
t.Fatal("query hooks must run on the OnTxBegin transaction")
|
||||||
|
}
|
||||||
|
if txs[OnTxBegin][0] == txs[OnTxBegin][1] || txs[BeforeResponse][0] != txs[OnTxBegin][1] {
|
||||||
|
t.Fatal("BeforeResponse must run on a second, distinct transaction")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOnTxBeginListBeforeResponseOnSecondTx(t *testing.T) {
|
||||||
|
var queries int
|
||||||
|
h := NewHandler(txFactory(&queries))
|
||||||
|
var txs []common.Database
|
||||||
|
h.Hooks().Register(OnTxBegin, func(c *HookContext) error { txs = append(txs, c.Tx); return nil })
|
||||||
|
var respTx common.Database
|
||||||
|
h.Hooks().Register(BeforeResponse, func(c *HookContext) error { respTx = c.Tx; return nil })
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
h.SqlQueryList("SELECT * FROM users", SqlQueryOptions{NoCount: true})(w, createTestRequest("GET", "/t", nil, nil, nil))
|
||||||
|
|
||||||
|
if w.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status %d body %s", w.Code, w.Body)
|
||||||
|
}
|
||||||
|
if len(txs) != 2 || txs[0] == txs[1] || respTx != txs[1] {
|
||||||
|
t.Fatalf("expected 2 distinct transactions with BeforeResponse on the second, got %d", len(txs))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOnTxBeginErrorAnswersTransactionError(t *testing.T) {
|
||||||
|
var queries int
|
||||||
|
h := NewHandler(txFactory(&queries))
|
||||||
|
h.Hooks().Register(OnTxBegin, func(*HookContext) error { return errors.New("secret detail") })
|
||||||
|
|
||||||
|
for name, run := range map[string]HTTPFuncType{
|
||||||
|
"single": h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{}),
|
||||||
|
"list": h.SqlQueryList("SELECT * FROM users", SqlQueryOptions{NoCount: true}),
|
||||||
|
} {
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
run(w, createTestRequest("GET", "/t", nil, nil, nil))
|
||||||
|
if w.Code != http.StatusInternalServerError {
|
||||||
|
t.Fatalf("%s: status %d body %s", name, w.Code, w.Body)
|
||||||
|
}
|
||||||
|
if strings.Contains(w.Body.String(), "secret detail") {
|
||||||
|
t.Fatalf("%s: hook error leaked to client: %s", name, w.Body)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if queries != 0 {
|
||||||
|
t.Fatalf("no query may run after a failed OnTxBegin, ran %d", queries)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user