diff --git a/audit/single_tran.md b/audit/single_tran.md index 0ec5e90..1039c5e 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -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 | | 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 | -| 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 | | ## 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`. - 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). -- 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 - Existing: per-spec `handler_test.go`, `hooks_test.go`, `integration_test.go`; models in `pkg/testmodels/business.go`; `dbtrace` unit tests. diff --git a/pkg/funcspec/function_api.go b/pkg/funcspec/function_api.go index 1c5f6ec..9afabe5 100644 --- a/pkg/funcspec/function_api.go +++ b/pkg/funcspec/function_api.go @@ -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) +} diff --git a/pkg/funcspec/hooks.go b/pkg/funcspec/hooks.go index e417b6c..de7e7d9 100644 --- a/pkg/funcspec/hooks.go +++ b/pkg/funcspec/hooks.go @@ -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 diff --git a/pkg/funcspec/tx_test.go b/pkg/funcspec/tx_test.go new file mode 100644 index 0000000..3ad946f --- /dev/null +++ b/pkg/funcspec/tx_test.go @@ -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) + } +}