diff --git a/audit/single_tran.md b/audit/single_tran.md index 2bbec0c..83dbe22 100644 --- a/audit/single_tran.md +++ b/audit/single_tran.md @@ -75,7 +75,7 @@ | 1 | DONE | Delete in one tx (single + batch, per-item hooks inside tx) | `resolvespec/handler.go`, `restheadspec/handler.go` | fixes 2 pool connections + race + RLS | | 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 | -| 4 | TODO | 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 | | 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 | | @@ -89,7 +89,9 @@ - DONE P3: restheadspec update re-fetch + `BeforeScan` + `AfterUpdate` and `AfterCreate` run in a second short `runInTx`; resolvespec update re-fetches (single, both batch paths) run in a second short `runInTx`. Fixed the pool reads inside the first tx (resolvespec single/batch update existing-record select, restheadspec update existence select) to use `tx`. Tests: `pkg/*/update_tx_test.go` (restheadspec uses the bun adapter; the pgsql adapter cannot build model-based updates). - NOTE: resolvespec fires no `AfterCreate`/`AfterRead`/`AfterUpdate`-post-commit hooks other than `AfterUpdate` inside the tx; nothing more to move there. - 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. -- NEXT: P4. +- 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"). +- NEXT: P5. ## 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/mqttspec/handler.go b/pkg/mqttspec/handler.go index 841a0c9..98185cf 100644 --- a/pkg/mqttspec/handler.go +++ b/pkg/mqttspec/handler.go @@ -3,6 +3,7 @@ package mqttspec import ( "context" "encoding/json" + "errors" "fmt" "reflect" "strings" @@ -313,42 +314,73 @@ func (h *Handler) handleRequest(client *Client, msg *Message) { } } -// handleRead processes a read operation +// stageError marks which stage of an operation failed inside a transaction so the +// right error response is sent once the transaction has rolled back. +type stageError struct { + code string + err error +} + +func (e *stageError) Error() string { return e.err.Error() } +func (e *stageError) Unwrap() error { return e.err } + +func hookStage(err error) error { return &stageError{code: "hook_error", err: err} } +func opStage(code string, err error) error { return &stageError{code: code, err: err} } + +// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin fired +// first. Transactions are per message, never per client connection. +func (h *Handler) runInTx(hookCtx *HookContext, body func(tx common.Database) error) error { + return common.RunRequestTx(hookCtx.Context, h.db, hookCtx, func() error { + return h.hooks.Execute(OnTxBegin, hookCtx) + }, body) +} + +// sendTxError sends the response for an error returned by runInTx. Errors that are +// not stage errors (begin, OnTxBegin, commit) carry no detail to the client. +func (h *Handler) sendTxError(client *Client, msgID string, err error) { + var stage *stageError + if errors.As(err, &stage) { + logger.Error("[MQTTSpec] %s: %v", stage.code, stage.err) + h.sendError(client.ID, msgID, stage.code, stage.err.Error()) + return + } + logger.Error("[MQTTSpec] Transaction failed: %v", err) + h.sendError(client.ID, msgID, "transaction_error", "Transaction failed") +} + +// handleRead processes a read operation; hooks and queries share one transaction. func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { - logger.Error("[MQTTSpec] BeforeRead hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) - return - } - - // Perform read operation - var data interface{} var metadata map[string]interface{} - var err error - if hookCtx.ID != "" { - // Read single record by ID - data, err = h.readByID(hookCtx) - metadata = map[string]interface{}{"total": 1} - } else { - // Read multiple records - data, metadata, err = h.readMultiple(hookCtx) - } + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { + return hookStage(err) + } + var data interface{} + var err error + if hookCtx.ID != "" { + // Read single record by ID + data, err = h.readByID(hookCtx) + metadata = map[string]interface{}{"total": 1} + } else { + // Read multiple records + data, metadata, err = h.readMultiple(hookCtx) + } + if err != nil { + return opStage("read_error", err) + } + + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) if err != nil { - logger.Error("[MQTTSpec] Read operation failed: %v", err) - h.sendError(client.ID, msg.ID, "read_error", err.Error()) - return - } - - // Update hook context - hookCtx.Result = data - - // Execute after hook - if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { - logger.Error("[MQTTSpec] AfterRead hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) + h.sendTxError(client, msg.ID, err) return } @@ -356,30 +388,49 @@ func (h *Handler) handleRead(client *Client, msg *Message, hookCtx *HookContext) h.sendResponse(client.ID, msg.ID, hookCtx.Result, metadata) } -// handleCreate processes a create operation +// handleCreate processes a create operation. The insert runs in the first +// transaction; the re-fetch (to capture DB defaults/triggers) and AfterCreate run +// in a second short transaction. func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { - logger.Error("[MQTTSpec] BeforeCreate hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) - return - } + var data interface{} - // Perform create operation - data, err := h.create(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { + return hookStage(err) + } + + var err error + data, err = h.create(hookCtx) + if err != nil { + return opStage("create_error", err) + } + return nil + }) if err != nil { - logger.Error("[MQTTSpec] Create operation failed: %v", err) - h.sendError(client.ID, msg.ID, "create_error", err.Error()) + h.sendTxError(client, msg.ID, err) return } - // Update hook context - hookCtx.Result = data + err = h.runInTx(hookCtx, func(tx common.Database) error { + if pkVal := reflection.GetPrimaryKeyValue(hookCtx.ModelPtr); pkVal != nil { + hookCtx.ID = fmt.Sprintf("%v", pkVal) + var err error + data, err = h.readByID(hookCtx) + if err != nil { + return opStage("create_error", err) + } + } - // Execute after hook - if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { - logger.Error("[MQTTSpec] AfterCreate hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(client, msg.ID, err) return } @@ -390,30 +441,42 @@ func (h *Handler) handleCreate(client *Client, msg *Message, hookCtx *HookContex h.notifySubscribers(hookCtx.Schema, hookCtx.Entity, OperationCreate, data) } -// handleUpdate processes an update operation +// handleUpdate processes an update operation. The update runs in the first +// transaction; the re-fetch and AfterUpdate run in a second short transaction. func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { - logger.Error("[MQTTSpec] BeforeUpdate hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) - return - } + var data interface{} - // Perform update operation - data, err := h.update(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { + return hookStage(err) + } + if err := h.update(hookCtx); err != nil { + return opStage("update_error", err) + } + return nil + }) if err != nil { - logger.Error("[MQTTSpec] Update operation failed: %v", err) - h.sendError(client.ID, msg.ID, "update_error", err.Error()) + h.sendTxError(client, msg.ID, err) return } - // Update hook context - hookCtx.Result = data + err = h.runInTx(hookCtx, func(tx common.Database) error { + var err error + data, err = h.readByID(hookCtx) + if err != nil { + return opStage("update_error", err) + } - // Execute after hook - if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { - logger.Error("[MQTTSpec] AfterUpdate hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(client, msg.ID, err) return } @@ -424,26 +487,22 @@ func (h *Handler) handleUpdate(client *Client, msg *Message, hookCtx *HookContex h.notifySubscribers(hookCtx.Schema, hookCtx.Entity, OperationUpdate, data) } -// handleDelete processes a delete operation +// handleDelete processes a delete operation; hooks and delete share one transaction. func (h *Handler) handleDelete(client *Client, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { - logger.Error("[MQTTSpec] BeforeDelete hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) - return - } - - // Perform delete operation - if err := h.delete(hookCtx); err != nil { - logger.Error("[MQTTSpec] Delete operation failed: %v", err) - h.sendError(client.ID, msg.ID, "delete_error", err.Error()) - return - } - - // Execute after hook - if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { - logger.Error("[MQTTSpec] AfterDelete hook failed: %v", err) - h.sendError(client.ID, msg.ID, "hook_error", err.Error()) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { + return hookStage(err) + } + if err := h.delete(hookCtx); err != nil { + return opStage("delete_error", err) + } + if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(client, msg.ID, err) return } @@ -671,7 +730,7 @@ func (h *Handler) getTableName(schema, entity string, model interface{}) string // readByID reads a single record by ID func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { - query := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Add ID filter pkName := reflection.GetPrimaryKeyName(hookCtx.Model) @@ -711,7 +770,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { // readMultiple reads multiple records func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata map[string]interface{}, err error) { - query := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Apply options if hookCtx.Options != nil { @@ -786,7 +845,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata // Get count metadata = make(map[string]interface{}) - countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + countQuery := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) if hookCtx.Options != nil { for _, filter := range hookCtx.Options.Filters { if cond, jargs, ok := common.BuildJSONFilterCondition(hookCtx.Model, "", filter.Column, filter.Operator, filter.Value); ok { @@ -835,22 +894,16 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) { } // Insert record - query := h.db.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) if _, err := query.Exec(hookCtx.Context); err != nil { return nil, fmt.Errorf("failed to create record: %w", err) } - // Re-fetch the created record to capture DB-generated defaults/triggers. - if pkVal := reflection.GetPrimaryKeyValue(hookCtx.ModelPtr); pkVal != nil { - hookCtx.ID = fmt.Sprintf("%v", pkVal) - return h.readByID(hookCtx) - } - return hookCtx.ModelPtr, nil } // update updates an existing record -func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { +func (h *Handler) update(hookCtx *HookContext) error { // Convert request data to a map var updates map[string]interface{} if m, ok := hookCtx.Data.(map[string]interface{}); ok { @@ -858,10 +911,10 @@ func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { } else { dataBytes, err := json.Marshal(hookCtx.Data) if err != nil { - return nil, fmt.Errorf("failed to marshal data: %w", err) + return fmt.Errorf("failed to marshal data: %w", err) } if err := json.Unmarshal(dataBytes, &updates); err != nil { - return nil, fmt.Errorf("failed to unmarshal data into map: %w", err) + return fmt.Errorf("failed to unmarshal data into map: %w", err) } } @@ -872,21 +925,20 @@ func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls) if len(values) > 0 { - query := h.db.NewUpdate().Table(hookCtx.TableName).SetMap(values). + query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID) if _, err := query.Exec(hookCtx.Context); err != nil { - return nil, fmt.Errorf("failed to update record: %w", err) + return fmt.Errorf("failed to update record: %w", err) } } - // Fetch updated record - return h.readByID(hookCtx) + return nil } // delete deletes a record func (h *Handler) delete(hookCtx *HookContext) error { - query := h.db.NewDelete().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewDelete().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Add ID filter pkName := reflection.GetPrimaryKeyName(hookCtx.Model) diff --git a/pkg/mqttspec/handler_test.go b/pkg/mqttspec/handler_test.go index db30ae2..6cc4bf0 100644 --- a/pkg/mqttspec/handler_test.go +++ b/pkg/mqttspec/handler_test.go @@ -773,7 +773,7 @@ func TestHandler_HandleIncomingMessage_ValidMessage(t *testing.T) { } func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { - newHook := func(id string, data map[string]interface{}) *HookContext { + newHook := func(handler *Handler, id string, data map[string]interface{}) *HookContext { return &HookContext{ Context: context.Background(), TableName: "users", @@ -784,6 +784,7 @@ func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { ID: id, Data: data, Options: &common.RequestOptions{}, + Tx: handler.db, } } seed := func(t *testing.T, db *gorm.DB) { @@ -794,7 +795,7 @@ func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { handler, db := setupTestHandler(t) seed(t, db) - _, err := handler.update(newHook("1", map[string]interface{}{"name": ""})) + err := handler.update(newHook(handler, "1", map[string]interface{}{"name": ""})) require.NoError(t, err) var got TestUser @@ -808,7 +809,7 @@ func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { handler, db := setupTestHandler(t) seed(t, db) - _, err := handler.update(newHook("1", map[string]interface{}{"status": "inactive"})) + err := handler.update(newHook(handler, "1", map[string]interface{}{"status": "inactive"})) require.NoError(t, err) var got TestUser @@ -823,7 +824,7 @@ func TestHandler_Update_OnlyPresentKeysChange(t *testing.T) { handler.SetDisallowNulls(true) seed(t, db) - _, err := handler.update(newHook("1", map[string]interface{}{"name": nil, "status": ""})) + err := handler.update(newHook(handler, "1", map[string]interface{}{"name": nil, "status": ""})) require.NoError(t, err) var got TestUser diff --git a/pkg/mqttspec/hooks.go b/pkg/mqttspec/hooks.go index 205be3c..3277441 100644 --- a/pkg/mqttspec/hooks.go +++ b/pkg/mqttspec/hooks.go @@ -50,6 +50,11 @@ const ( // BeforeOp fires immediately before every SQL operation (read, create, update, delete) BeforeOp = websocketspec.BeforeOp + + // OnTxBegin fires once, first, inside every transaction the handler opens for a + // read/create/update/delete message (including the second short transaction for + // post-commit work). hookCtx.Tx is the transaction. + OnTxBegin = websocketspec.OnTxBegin ) // NewHookRegistry creates a new hook registry diff --git a/pkg/mqttspec/tx_test.go b/pkg/mqttspec/tx_test.go new file mode 100644 index 0000000..4293e92 --- /dev/null +++ b/pkg/mqttspec/tx_test.go @@ -0,0 +1,91 @@ +package mqttspec + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func newTxHook(handler *Handler, data interface{}) *HookContext { + return &HookContext{ + Context: context.Background(), + TableName: "users", + Model: &TestUser{}, + ModelPtr: &TestUser{}, + Schema: "public", + Entity: "users", + ID: "1", + Data: data, + Options: &common.RequestOptions{}, + Metadata: map[string]interface{}{}, + Tx: handler.db, + } +} + +func TestHandler_OnTxBeginFiresFirstOnTransaction(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + + var order []string + var txs []common.Database + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { + order = append(order, "begin") + txs = append(txs, c.Tx) + return nil + }) + handler.hooks.Register(BeforeDelete, func(c *HookContext) error { + order = append(order, "before") + return nil + }) + + handler.handleDelete(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, nil)) + + var count int64 + require.NoError(t, db.Model(&TestUser{}).Where("id = 1").Count(&count).Error) + assert.Zero(t, count) + require.Len(t, txs, 1) + assert.NotEqual(t, handler.db, txs[0]) + assert.Equal(t, []string{"begin", "before"}, order) +} + +func TestHandler_OnTxBeginErrorAbortsWithoutWriting(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { return errors.New("no user") }) + + handler.handleDelete(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, nil)) + + var count int64 + require.NoError(t, db.Model(&TestUser{}).Where("id = 1").Count(&count).Error) + assert.Equal(t, int64(1), count) +} + +func TestHandler_UpdateRunsAfterHookOnSecondTransaction(t *testing.T) { + handler, db := setupTestHandler(t) + require.NoError(t, db.Create(&TestUser{ID: 1, Name: "a", Email: "a@example.com", Status: "active"}).Error) + + var begins []common.Database + var afterTx common.Database + handler.hooks.Register(OnTxBegin, func(c *HookContext) error { + begins = append(begins, c.Tx) + return nil + }) + handler.hooks.Register(AfterUpdate, func(c *HookContext) error { + afterTx = c.Tx + return nil + }) + + handler.handleUpdate(&Client{ID: "c1"}, &Message{ID: "m1"}, newTxHook(handler, map[string]interface{}{"name": "b"})) + + var got TestUser + require.NoError(t, db.First(&got, 1).Error) + assert.Equal(t, "b", got.Name) + require.Len(t, begins, 2) + assert.NotEqual(t, begins[0], begins[1]) + assert.Equal(t, begins[1], afterTx) +} diff --git a/pkg/websocketspec/handler.go b/pkg/websocketspec/handler.go index 0dbe2f4..352d1d8 100644 --- a/pkg/websocketspec/handler.go +++ b/pkg/websocketspec/handler.go @@ -221,49 +221,84 @@ func (h *Handler) handleRequest(conn *Connection, msg *Message) { } } -// handleRead processes a read operation +// stageError marks which stage of an operation failed inside a transaction so the +// right error response is sent once the transaction has rolled back. +type stageError struct { + code string + hook bool + err error +} + +func (e *stageError) Error() string { return e.err.Error() } +func (e *stageError) Unwrap() error { return e.err } + +func hookStage(err error) error { return &stageError{code: "hook_error", hook: true, err: err} } +func opStage(code string, err error) error { + return &stageError{code: code, err: err} +} + +// runInTx runs body in a transaction with hookCtx.Tx set to it and OnTxBegin fired +// first. Transactions are per message, never per connection. +func (h *Handler) runInTx(hookCtx *HookContext, body func(tx common.Database) error) error { + return common.RunRequestTx(hookCtx.Context, h.db, hookCtx, func() error { + return h.hooks.Execute(OnTxBegin, hookCtx) + }, body) +} + +// sendTxError sends the response for an error returned by runInTx. Errors that are +// not stage errors (begin, OnTxBegin, commit) carry no detail to the client. +func (h *Handler) sendTxError(conn *Connection, msgID string, err error) { + var stage *stageError + switch { + case errors.As(err, &stage) && stage.hook: + logger.Error("[WebSocketSpec] %s: %v", stage.code, stage.err) + _ = conn.SendJSON(NewErrorResponse(msgID, stage.code, stage.err.Error())) + case errors.As(err, &stage): + logger.Error("[WebSocketSpec] %s: %v", stage.code, stage.err) + _ = conn.SendJSON(newErrorResponseFromErr(msgID, stage.code, stage.err)) + default: + logger.Error("[WebSocketSpec] Transaction failed: %v", err) + _ = conn.SendJSON(NewErrorResponse(msgID, "transaction_error", "Transaction failed")) + } +} + +// handleRead processes a read operation; hooks and queries share one transaction. func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { - logger.Error("[WebSocketSpec] BeforeRead hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) - return - } - - // Perform read operation - var data interface{} var metadata map[string]interface{} - var err error - // Check if FetchRowNumber is specified (treat as single record read) - isFetchRowNumber := hookCtx.Options != nil && hookCtx.Options.FetchRowNumber != nil && *hookCtx.Options.FetchRowNumber != "" + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeRead, hookCtx); err != nil { + return hookStage(err) + } - if hookCtx.ID != "" || isFetchRowNumber { - // Read single record by ID or FetchRowNumber - data, err = h.readByID(hookCtx) - metadata = map[string]interface{}{"total": 1} - // The row number is already set on the record itself via setRowNumbersOnRecords - } else { - // Read multiple records - data, metadata, err = h.readMultiple(hookCtx) - } + var data interface{} + var err error + // Check if FetchRowNumber is specified (treat as single record read) + isFetchRowNumber := hookCtx.Options != nil && hookCtx.Options.FetchRowNumber != nil && *hookCtx.Options.FetchRowNumber != "" + if hookCtx.ID != "" || isFetchRowNumber { + // Read single record by ID or FetchRowNumber + data, err = h.readByID(hookCtx) + metadata = map[string]interface{}{"total": 1} + // The row number is already set on the record itself via setRowNumbersOnRecords + } else { + // Read multiple records + data, metadata, err = h.readMultiple(hookCtx) + } + if err != nil { + return opStage("read_error", err) + } + + // Update hook context with result + hookCtx.Result = data + + if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) if err != nil { - logger.Error("[WebSocketSpec] Read operation failed: %v", err) - errResp := newErrorResponseFromErr(msg.ID, "read_error", err) - _ = conn.SendJSON(errResp) - return - } - - // Update hook context with result - hookCtx.Result = data - - // Execute after hook - if err := h.hooks.Execute(AfterRead, hookCtx); err != nil { - logger.Error("[WebSocketSpec] AfterRead hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) + h.sendTxError(conn, msg.ID, err) return } @@ -273,33 +308,49 @@ func (h *Handler) handleRead(conn *Connection, msg *Message, hookCtx *HookContex _ = conn.SendJSON(resp) } -// handleCreate processes a create operation +// handleCreate processes a create operation. The insert runs in the first +// transaction; the re-fetch (to capture DB defaults/triggers) and AfterCreate run +// in a second short transaction. func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { - logger.Error("[WebSocketSpec] BeforeCreate hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) - return - } + var data interface{} - // Perform create operation - data, err := h.create(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeCreate, hookCtx); err != nil { + return hookStage(err) + } + + var err error + data, err = h.create(hookCtx) + if err != nil { + return opStage("create_error", err) + } + return nil + }) if err != nil { - logger.Error("[WebSocketSpec] Create operation failed: %v", err) - errResp := newErrorResponseFromErr(msg.ID, "create_error", err) - _ = conn.SendJSON(errResp) + h.sendTxError(conn, msg.ID, err) return } - // Update hook context - hookCtx.Result = data + err = h.runInTx(hookCtx, func(tx common.Database) error { + if reflection.GetPrimaryKeyValue(hookCtx.ModelPtr) != nil { + hookCtx.ID = fmt.Sprintf("%v", reflection.GetPrimaryKeyValue(hookCtx.ModelPtr)) + var err error + data, err = h.readByID(hookCtx) + if err != nil { + return opStage("create_error", err) + } + } - // Execute after hook - if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { - logger.Error("[WebSocketSpec] AfterCreate hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterCreate, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(conn, msg.ID, err) return } @@ -311,33 +362,42 @@ func (h *Handler) handleCreate(conn *Connection, msg *Message, hookCtx *HookCont h.notifySubscribers(hookCtx.Schema, hookCtx.Entity, OperationCreate, data) } -// handleUpdate processes an update operation +// handleUpdate processes an update operation. The update runs in the first +// transaction; the re-fetch and AfterUpdate run in a second short transaction. func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { - logger.Error("[WebSocketSpec] BeforeUpdate hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) - return - } + var data interface{} - // Perform update operation - data, err := h.update(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeUpdate, hookCtx); err != nil { + return hookStage(err) + } + if err := h.update(hookCtx); err != nil { + return opStage("update_error", err) + } + return nil + }) if err != nil { - logger.Error("[WebSocketSpec] Update operation failed: %v", err) - errResp := newErrorResponseFromErr(msg.ID, "update_error", err) - _ = conn.SendJSON(errResp) + h.sendTxError(conn, msg.ID, err) return } - // Update hook context - hookCtx.Result = data + err = h.runInTx(hookCtx, func(tx common.Database) error { + var err error + data, err = h.readByID(hookCtx) + if err != nil { + return opStage("update_error", err) + } - // Execute after hook - if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { - logger.Error("[WebSocketSpec] AfterUpdate hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) + // Update hook context + hookCtx.Result = data + + if err := h.hooks.Execute(AfterUpdate, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) + if err != nil { + h.sendTxError(conn, msg.ID, err) return } @@ -349,30 +409,22 @@ func (h *Handler) handleUpdate(conn *Connection, msg *Message, hookCtx *HookCont h.notifySubscribers(hookCtx.Schema, hookCtx.Entity, OperationUpdate, data) } -// handleDelete processes a delete operation +// handleDelete processes a delete operation; hooks and delete share one transaction. func (h *Handler) handleDelete(conn *Connection, msg *Message, hookCtx *HookContext) { - // Execute before hook - if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { - logger.Error("[WebSocketSpec] BeforeDelete hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) - return - } - - // Perform delete operation - err := h.delete(hookCtx) + err := h.runInTx(hookCtx, func(tx common.Database) error { + if err := h.hooks.ExecuteBeforeOp(BeforeDelete, hookCtx); err != nil { + return hookStage(err) + } + if err := h.delete(hookCtx); err != nil { + return opStage("delete_error", err) + } + if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { + return hookStage(err) + } + return nil + }) if err != nil { - logger.Error("[WebSocketSpec] Delete operation failed: %v", err) - errResp := newErrorResponseFromErr(msg.ID, "delete_error", err) - _ = conn.SendJSON(errResp) - return - } - - // Execute after hook - if err := h.hooks.Execute(AfterDelete, hookCtx); err != nil { - logger.Error("[WebSocketSpec] AfterDelete hook failed: %v", err) - errResp := NewErrorResponse(msg.ID, "hook_error", err.Error()) - _ = conn.SendJSON(errResp) + h.sendTxError(conn, msg.ID, err) return } @@ -548,7 +600,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { fetchRowNumberPKValue := *hookCtx.Options.FetchRowNumber logger.Debug("[WebSocketSpec] FetchRowNumber: Fetching row number for PK %s = %s", pkName, fetchRowNumberPKValue) - rowNum, err := h.FetchRowNumber(hookCtx.Context, hookCtx.TableName, pkName, fetchRowNumberPKValue, hookCtx.Options, hookCtx.Model) + rowNum, err := h.fetchRowNumber(hookCtx.Context, hookCtx.Tx, hookCtx.TableName, pkName, fetchRowNumberPKValue, hookCtx.Options, hookCtx.Model) if err != nil { return nil, fmt.Errorf("failed to fetch row number: %w", err) } @@ -560,7 +612,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { hookCtx.ID = fetchRowNumberPKValue } - query := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Add ID filter query = query.Where(fmt.Sprintf("%s = ?", pkName), hookCtx.ID) @@ -604,7 +656,7 @@ func (h *Handler) readByID(hookCtx *HookContext) (interface{}, error) { } func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata map[string]interface{}, err error) { - query := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Apply options (simplified implementation) if hookCtx.Options != nil { @@ -669,7 +721,7 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata // Get count metadata = make(map[string]interface{}) - countQuery := h.db.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + countQuery := hookCtx.Tx.NewSelect().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) if hookCtx.Options != nil { for _, filter := range hookCtx.Options.Filters { cond, args := h.buildFilterCondition(filter, hookCtx.Model) @@ -705,21 +757,15 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) { } // Insert record - query := h.db.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) if _, err := query.Exec(hookCtx.Context); err != nil { return nil, fmt.Errorf("failed to create record: %w", err) } - // Re-fetch the created record to capture DB-generated defaults/triggers. - if pkVal := reflection.GetPrimaryKeyValue(hookCtx.ModelPtr); pkVal != nil { - hookCtx.ID = fmt.Sprintf("%v", pkVal) - return h.readByID(hookCtx) - } - return hookCtx.ModelPtr, nil } -func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { +func (h *Handler) update(hookCtx *HookContext) error { // Convert request data to a map var updates map[string]interface{} if m, ok := hookCtx.Data.(map[string]interface{}); ok { @@ -727,10 +773,10 @@ func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { } else { dataBytes, err := json.Marshal(hookCtx.Data) if err != nil { - return nil, fmt.Errorf("failed to marshal data: %w", err) + return fmt.Errorf("failed to marshal data: %w", err) } if err := json.Unmarshal(dataBytes, &updates); err != nil { - return nil, fmt.Errorf("failed to unmarshal data into map: %w", err) + return fmt.Errorf("failed to unmarshal data into map: %w", err) } } @@ -741,20 +787,19 @@ func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls) if len(values) > 0 { - query := h.db.NewUpdate().Table(hookCtx.TableName).SetMap(values). + query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID) if _, err := query.Exec(hookCtx.Context); err != nil { - return nil, fmt.Errorf("failed to update record: %w", err) + return fmt.Errorf("failed to update record: %w", err) } } - // Fetch updated record - return h.readByID(hookCtx) + return nil } func (h *Handler) delete(hookCtx *HookContext) error { - query := h.db.NewDelete().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + query := hookCtx.Tx.NewDelete().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) // Add ID filter pkName := reflection.GetPrimaryKeyName(hookCtx.Model) @@ -966,6 +1011,11 @@ func (h *Handler) getOperatorSQL(operator string) string { // FetchRowNumber calculates the row number of a specific record based on sorting and filtering // Returns the 1-based row number of the record with the given primary key value func (h *Handler) FetchRowNumber(ctx context.Context, tableName string, pkName string, pkValue string, options *common.RequestOptions, model interface{}) (int64, error) { + return h.fetchRowNumber(ctx, h.db, tableName, pkName, pkValue, options, model) +} + +// fetchRowNumber is FetchRowNumber on the given database or transaction. +func (h *Handler) fetchRowNumber(ctx context.Context, db common.Database, tableName string, pkName string, pkValue string, options *common.RequestOptions, model interface{}) (int64, error) { defer func() { if r := recover(); r != nil { logger.Error("[WebSocketSpec] Panic during FetchRowNumber: %v", r) @@ -1033,7 +1083,7 @@ func (h *Handler) FetchRowNumber(ctx context.Context, tableName string, pkName s var result []struct { RN int64 `bun:"rn"` } - err := h.db.Query(ctx, &result, queryStr, whereArgs...) + err := db.Query(ctx, &result, queryStr, whereArgs...) if err != nil { return 0, fmt.Errorf("failed to fetch row number: %w", err) } diff --git a/pkg/websocketspec/hooks.go b/pkg/websocketspec/hooks.go index 171470d..5d4e55c 100644 --- a/pkg/websocketspec/hooks.go +++ b/pkg/websocketspec/hooks.go @@ -64,6 +64,13 @@ const ( // Unlike BeforeHandle, which fires once before operation dispatch, BeforeOp 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 for a + // read/create/update/delete message (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 context information for hook execution @@ -128,6 +135,9 @@ type HookContext struct { Metadata map[string]interface{} } +// SetTx points the context at the transaction in use (common.TxContext). +func (c *HookContext) SetTx(tx common.Database) { c.Tx = tx } + // HookFunc is a function that processes a hook type HookFunc func(*HookContext) error diff --git a/pkg/websocketspec/tx_test.go b/pkg/websocketspec/tx_test.go new file mode 100644 index 0000000..12e410d --- /dev/null +++ b/pkg/websocketspec/tx_test.go @@ -0,0 +1,173 @@ +package websocketspec + +import ( + "context" + "encoding/json" + "errors" + "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, *Connection, *HookContext) { + t.Helper() + db, mock, err := sqlmock.New() + if err != nil { + t.Fatal(err) + } + // One connection: any statement that bypasses the transaction while it is + // open 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()) + conn := NewConnection("c1", nil, h) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + t.Cleanup(cancel) + hookCtx := &HookContext{ + Context: ctx, + Handler: h, + Schema: "public", + Entity: "items", + TableName: "public.items", + Model: txItem{}, + ModelPtr: &txItem{}, + ID: "7", + Tx: h.db, + Metadata: map[string]interface{}{}, + } + return h, mock, conn, hookCtx +} + +func response(t *testing.T, conn *Connection) ResponseMessage { + t.Helper() + select { + case raw := <-conn.send: + var resp ResponseMessage + if err := json.Unmarshal(raw, &resp); err != nil { + t.Fatal(err) + } + return resp + default: + t.Fatal("no response sent") + return ResponseMessage{} + } +} + +func TestDeleteRunsHooksAndDeleteInOneTransaction(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + var order []string + var txs []common.Database + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { + order = append(order, "begin") + txs = append(txs, c.Tx) + return nil + }) + h.Hooks().Register(BeforeDelete, func(c *HookContext) error { + order = append(order, "before") + return nil + }) + h.Hooks().Register(AfterDelete, func(c *HookContext) error { + order = append(order, "after") + return nil + }) + + mock.ExpectBegin() + mock.ExpectExec(`DELETE FROM`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + + h.handleDelete(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if !response(t, conn).Success { + t.Fatal("expected success") + } + if len(txs) != 1 || txs[0] == h.db || len(order) != 3 || order[0] != "begin" { + t.Fatalf("OnTxBegin must fire once, first, on the transaction: %v", order) + } +} + +func TestUpdateRefetchAndAfterHookRunInSecondTransaction(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + hookCtx.Data = map[string]interface{}{"name": "b"} + var begins []common.Database + var afterTx common.Database + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { + begins = append(begins, c.Tx) + return nil + }) + h.Hooks().Register(AfterUpdate, func(c *HookContext) error { + afterTx = c.Tx + return nil + }) + + mock.ExpectBegin() + mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectCommit() + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "b")) + mock.ExpectCommit() + + h.handleUpdate(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if !response(t, conn).Success { + t.Fatal("expected success") + } + if len(begins) != 2 || begins[0] == begins[1] || afterTx != begins[1] { + t.Fatalf("OnTxBegin must fire per transaction and AfterUpdate run on the second, got %d", len(begins)) + } +} + +func TestReadRunsInOneTransaction(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + var beforeTx common.Database + h.Hooks().Register(BeforeRead, func(c *HookContext) error { + beforeTx = c.Tx + return nil + }) + + mock.ExpectBegin() + mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a")) + mock.ExpectCommit() + + h.handleRead(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + if !response(t, conn).Success || beforeTx == nil || beforeTx == h.db { + t.Fatal("BeforeRead must run on the transaction") + } +} + +func TestOnTxBeginErrorRollsBackWithoutDetail(t *testing.T) { + h, mock, conn, hookCtx := newTxHarness(t) + h.Hooks().Register(OnTxBegin, func(c *HookContext) error { return errors.New("secret detail") }) + + mock.ExpectBegin() + mock.ExpectRollback() + + h.handleDelete(conn, &Message{ID: "m1"}, hookCtx) + + if err := mock.ExpectationsWereMet(); err != nil { + t.Fatal(err) + } + resp := response(t, conn) + if resp.Success || resp.Error == nil || resp.Error.Code != "transaction_error" || resp.Error.Message != "Transaction failed" { + t.Fatalf("unexpected response %+v", resp) + } +}