feat(tx): run websocketspec and mqttspec operations in per-message transactions

Reads and deletes run in one transaction; create and update write in one
and re-fetch plus After hooks in a second. OnTxBegin fires first in each.
This commit is contained in:
2026-09-30 22:46:46 +02:00
parent 97bcb44fdc
commit ed457eb14a
8 changed files with 611 additions and 227 deletions
+169 -119
View File
@@ -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)
}