feat(tx): run funcspec OnTxBegin and BeforeResponse in transactions

This commit is contained in:
2026-09-30 22:50:41 +02:00
parent 4cbe4f597d
commit 6b6f540ab0
4 changed files with 321 additions and 183 deletions
+207 -181
View File
@@ -192,131 +192,140 @@ func (h *Handler) SqlQueryList(sqlquery string, options SqlQueryOptions) HTTPFun
hookCtx.InputVars = inputvars
// Execute query within transaction
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
// Set transaction in hook context for hooks to use
hookCtx.Tx = tx
// bodyRan/bodyFailed tell a begin/OnTxBegin/commit failure (no response sent
// yet) from a body failure (sendError already answered).
var bodyRan, bodyFailed bool
err := h.runInTx(ctx, hookCtx, func(tx common.Database) error {
bodyRan = true
berr := func() error {
// Execute BeforeQueryList hook (inside transaction)
if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil {
logger.Error("BeforeQueryList hook failed: %v", 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)
// Execute BeforeQueryList hook (inside transaction)
if err := h.hooks.ExecuteBeforeOp(BeforeQueryList, hookCtx); err != nil {
logger.Error("BeforeQueryList hook failed: %v", err)
sendError(w, http.StatusBadRequest, "hook_error", "Hook execution failed", 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
// 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)
}
// 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
}
// Use potentially modified SQL query from hook
sqlquery = hookCtx.SQLQuery
sqlqueryCnt := sqlquery
// Normalize PostgreSQL types for proper JSON marshaling
dbobjlist = normalizePostgresTypesList(rows)
// Parse sorting and pagination parameters
sortcols, limit, offset := h.parsePaginationParams(r)
if options.NoCount {
total = int64(len(dbobjlist))
}
// 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
}
// 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
hookCtx.SortColumns = sortcols
hookCtx.Limit = limit
hookCtx.Offset = offset
fromPos := strings.Index(strings.ToLower(sqlquery), "from ")
orderbyPos := strings.Index(strings.ToLower(sqlquery), "order by")
// 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
if len(sortcols) > 0 && (orderbyPos < 0 || (orderbyPos > 0 && orderbyPos < fromPos)) {
sqlquery = fmt.Sprintf("%s \nORDER BY %s", sqlquery, ValidSQL(sortcols, "select"))
}
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 {
logger.Error("Transaction failed: %v", err)
if !bodyRan || !bodyFailed {
sendError(w, http.StatusInternalServerError, "transaction_error", "Transaction failed", nil)
}
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))
logger.Info("Serving: Records %d of %d", len(dbobjlist), total)
// Execute BeforeResponse hook. The transaction has already committed by
// this point, so hooks must use the pooled connection rather than the
// now-dead tx.
hookCtx.Tx = h.db
// Execute BeforeResponse hook in a second short transaction: the main one
// has already committed, and hooks must never get the pooled connection.
hookCtx.Result = dbobjlist
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)
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
return
@@ -558,88 +567,97 @@ func (h *Handler) SqlQuery(sqlquery string, options SqlQueryOptions) HTTPFuncTyp
hookCtx.InputVars = inputvars
// Execute query within transaction
err := h.db.RunInTransaction(ctx, func(tx common.Database) error {
// Set transaction in hook context for hooks to use
hookCtx.Tx = tx
// bodyRan/bodyFailed tell a begin/OnTxBegin/commit failure (no response sent
// yet) from a body failure (sendError already answered).
var bodyRan, bodyFailed bool
err := h.runInTx(ctx, hookCtx, func(tx common.Database) error {
bodyRan = true
berr := func() error {
// Execute BeforeQuery hook (inside transaction)
if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil {
logger.Error("BeforeQuery hook failed: %v", 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
// Execute BeforeQuery hook (inside transaction)
if err := h.hooks.ExecuteBeforeOp(BeforeQuery, hookCtx); err != nil {
logger.Error("BeforeQuery hook failed: %v", err)
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
// 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)
}
// Execute BeforeSQLExec hook
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
// 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
}
// Execute BeforeSQLExec hook
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
if len(rows) > 0 {
dbobj = normalizePostgresTypes(rows[0])
}
// 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
}
// Execute AfterSQLExec hook
hookCtx.Result = dbobj
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
}
if len(rows) > 0 {
dbobj = normalizePostgresTypes(rows[0])
}
// 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
}
// Execute AfterSQLExec hook
hookCtx.Result = dbobj
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
}
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 {
logger.Error("Transaction failed: %v", err)
if !bodyRan || !bodyFailed {
sendError(w, http.StatusInternalServerError, "transaction_error", "Transaction failed", nil)
}
return
}
// Execute BeforeResponse hook. The transaction has already committed by
// this point, so hooks must use the pooled connection rather than the
// now-dead tx.
hookCtx.Tx = h.db
// Execute BeforeResponse hook in a second short transaction: the main one
// has already committed, and hooks must never get the pooled connection.
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)
sendError(w, http.StatusInternalServerError, "hook_error", "Hook execution failed", err)
return
@@ -1249,3 +1267,11 @@ func normalizePostgresValue(value interface{}) interface{} {
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)
}
+9
View File
@@ -32,6 +32,12 @@ const (
// 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.
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
@@ -75,6 +81,9 @@ type HookContext struct {
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
// It receives a HookContext and can modify it or return an error
// If an error is returned, the operation will be aborted
+102
View File
@@ -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)
}
}