From 6ed2803158bf03b1c0f2dcbe682dee2b71eda19c Mon Sep 17 00:00:00 2001 From: what Date: Thu, 20 Aug 2026 17:08:29 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20Hooks=20=E6=96=B0=E5=A2=9E=20UseTx?= =?UTF-8?q?=EF=BC=8CBefore/After=20=E6=94=B9=E7=94=A8=E7=B1=BB=E5=9E=8B?= =?UTF-8?q?=E5=8C=96=E7=9A=84=20HookError?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 之前 Hooks.Before/After 返回裸 error,调用方没法区分"钩子本身失败"和"SQL 执行失败"。改成返回 *HookError(包一层 exec.HookError),配合 errors.As 能精确识别。 同时给 Hooks 加一个 UseTx(dataset) (QueryFactory, error) 方法,在 Before 之前调用,允许钩子在写操作真正执行前替换掉默认连接(比如需要自动开事务包住 写操作和写后回调的场景),返回 nil, nil 表示沿用默认连接。目前只有 InsertDataset/UpdateDataset/DeleteDataset 会调用 UseTx,SelectDataset (只读)不需要。 --- delete_dataset.go | 14 +++- exec/query_hooks.go | 23 +++++- exec/query_hooks_test.go | 149 +++++++++++++++++++++++++++++++++++++++ hook_error.go | 7 ++ insert_dataset.go | 14 +++- select_dataset.go | 9 ++- update_dataset.go | 14 +++- 7 files changed, 220 insertions(+), 10 deletions(-) create mode 100644 exec/query_hooks_test.go create mode 100644 hook_error.go diff --git a/delete_dataset.go b/delete_dataset.go index d7e1b0f..a3e33f6 100644 --- a/delete_dataset.go +++ b/delete_dataset.go @@ -238,12 +238,22 @@ func (dd *DeleteDataset) ReturnsColumns() bool { // See Dataset#ToUpdateSQL for arguments func (dd *DeleteDataset) Executor() (executor exec.QueryExecutor) { if dd.hooks != nil { - dd.SetError(dd.hooks.Before(dd)) + if qf, err := dd.hooks.UseTx(dd); err != nil { + dd.SetError(err) + } else if qf != nil { + dd.queryFactory = qf + } + if herr := dd.hooks.Before(dd); herr != nil { + dd.SetError(herr) + } } executor = dd.queryFactory.FromSQLBuilder(dd.deleteSQLBuilder()) if dd.hooks != nil { executor.Hook(func(result interface{}) error { - return dd.hooks.After(dd, result) + if herr := dd.hooks.After(dd, result); herr != nil { + return herr + } + return nil }) } return executor diff --git a/exec/query_hooks.go b/exec/query_hooks.go index a8812d3..1bfcaab 100644 --- a/exec/query_hooks.go +++ b/exec/query_hooks.go @@ -1,7 +1,26 @@ package exec +import "fmt" + +// HookError 是 Hooks.Before/After 失败时必须返回的类型,外部通过 db.HookError 引用。 +type HookError struct { + Err error +} + +func (e *HookError) Error() string { + return fmt.Sprintf("hook error: %s", e.Err) +} + +func (e *HookError) Unwrap() error { + return e.Err +} + // Hooks 钩子实例 type Hooks interface { - Before(dataset interface{}) error - After(dataset interface{}, result interface{}) error + Before(dataset interface{}) *HookError + After(dataset interface{}, result interface{}) *HookError + // UseTx 在 Before 之前调用,dataset 是即将执行的 *InsertDataset/*UpdateDataset/*DeleteDataset + // (类型与 Before 收到的一致),此时数据集尚未绑定要执行的连接。返回非 nil 的 QueryFactory 时, + // 调用方会用它替换默认连接来执行这次写操作;返回 nil, nil 表示沿用默认连接。 + UseTx(dataset interface{}) (QueryFactory, error) } diff --git a/exec/query_hooks_test.go b/exec/query_hooks_test.go new file mode 100644 index 0000000..5a3b52e --- /dev/null +++ b/exec/query_hooks_test.go @@ -0,0 +1,149 @@ +package exec_test + +import ( + "errors" + "testing" + + dbv2 "git.fsdpf.net/go/db" + "git.fsdpf.net/go/db/exec" + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/suite" +) + +// mockHooks 是 exec.Hooks 的最小实现,用来验证 UseTx/Before/After 与 db.HookError 的接线是否正确 +type mockHooks struct { + useTxFactory exec.QueryFactory + useTxErr error + beforeErr error + afterErr error + + useTxCalls int + beforeCalls int + afterCalls int +} + +func (h *mockHooks) UseTx(dataset interface{}) (exec.QueryFactory, error) { + h.useTxCalls++ + return h.useTxFactory, h.useTxErr +} + +func (h *mockHooks) Before(dataset interface{}) *exec.HookError { + h.beforeCalls++ + if h.beforeErr != nil { + return &exec.HookError{Err: h.beforeErr} + } + return nil +} + +func (h *mockHooks) After(dataset interface{}, result interface{}) *exec.HookError { + h.afterCalls++ + if h.afterErr != nil { + return &exec.HookError{Err: h.afterErr} + } + return nil +} + +type hooksSuite struct { + suite.Suite +} + +// TestBeforeError:Before 失败时 SQL 不会被执行,error 能用 errors.As 解出 db.HookError +func (hs *hooksSuite) TestBeforeError() { + mDB, mock, err := sqlmock.New() + hs.NoError(err) + + hooks := &mockHooks{beforeErr: errors.New("before failed")} + + ds := dbv2.New("mock", mDB).Insert("items"). + Rows(dbv2.Record{"name": "Test1"}). + WithHook(hooks) + + _, err = ds.Executor().Exec() + hs.Error(err) + + var hookErr *dbv2.HookError + hs.True(errors.As(err, &hookErr)) + hs.Equal("before failed", hookErr.Unwrap().Error()) + + hs.Equal(1, hooks.beforeCalls) + hs.Equal(0, hooks.afterCalls) + hs.NoError(mock.ExpectationsWereMet()) +} + +// TestAfterError:SQL 已经执行成功,After 失败时 error 依然能用 errors.As 解出 db.HookError, +// 且 result 仍然可用(区别于 SQL 本身执行失败) +func (hs *hooksSuite) TestAfterError() { + mDB, mock, err := sqlmock.New() + hs.NoError(err) + + mock.ExpectExec(`INSERT INTO "items"`).WillReturnResult(sqlmock.NewResult(1, 1)) + + hooks := &mockHooks{afterErr: errors.New("after failed")} + + ds := dbv2.New("mock", mDB).Insert("items"). + Rows(dbv2.Record{"name": "Test1"}). + WithHook(hooks) + + result, err := ds.Executor().Exec() + hs.Error(err) + hs.NotNil(result) + + rows, rowsErr := result.RowsAffected() + hs.NoError(rowsErr) + hs.EqualValues(1, rows) + + var hookErr *dbv2.HookError + hs.True(errors.As(err, &hookErr)) + hs.Equal("after failed", hookErr.Unwrap().Error()) + + hs.NoError(mock.ExpectationsWereMet()) +} + +// TestUseTxSwapsConnection:UseTx 返回的 QueryFactory 会替换默认连接,SQL 应该在返回的 +// factory 上执行,原始连接完全不会被用到 +func (hs *hooksSuite) TestUseTxSwapsConnection() { + origDB, origMock, err := sqlmock.New() + hs.NoError(err) + + txDB, txMock, err := sqlmock.New() + hs.NoError(err) + txMock.ExpectExec(`INSERT INTO "items"`).WillReturnResult(sqlmock.NewResult(1, 1)) + + hooks := &mockHooks{useTxFactory: exec.NewQueryFactory(txDB)} + + ds := dbv2.New("mock", origDB).Insert("items"). + Rows(dbv2.Record{"name": "Test1"}). + WithHook(hooks) + + result, err := ds.Executor().Exec() + hs.NoError(err) + hs.NotNil(result) + + hs.Equal(1, hooks.useTxCalls) + hs.NoError(txMock.ExpectationsWereMet()) + hs.NoError(origMock.ExpectationsWereMet()) +} + +// TestUseTxNil:UseTx 返回 nil factory 时沿用默认连接 +func (hs *hooksSuite) TestUseTxNil() { + mDB, mock, err := sqlmock.New() + hs.NoError(err) + + mock.ExpectExec(`INSERT INTO "items"`).WillReturnResult(sqlmock.NewResult(1, 1)) + + hooks := &mockHooks{} + + ds := dbv2.New("mock", mDB).Insert("items"). + Rows(dbv2.Record{"name": "Test1"}). + WithHook(hooks) + + _, err = ds.Executor().Exec() + hs.NoError(err) + + hs.Equal(1, hooks.useTxCalls) + hs.NoError(mock.ExpectationsWereMet()) +} + +func TestHooks(t *testing.T) { + suite.Run(t, new(hooksSuite)) +} diff --git a/hook_error.go b/hook_error.go new file mode 100644 index 0000000..2cdde17 --- /dev/null +++ b/hook_error.go @@ -0,0 +1,7 @@ +package db + +import "git.fsdpf.net/go/db/exec" + +// HookError 是 Hooks.Before/After 失败时的包装错误:Before 失败时 SQL 还未执行,After 失败时 +// SQL 语句本身已经执行成功、失败的是钩子本身。调用方可以用 errors.As 把它和真正的 SQL 执行失败区分开。 +type HookError = exec.HookError diff --git a/insert_dataset.go b/insert_dataset.go index 02cee76..d003daf 100644 --- a/insert_dataset.go +++ b/insert_dataset.go @@ -306,12 +306,22 @@ func (id *InsertDataset) ReturnsColumns() bool { // db.Insert("test").Rows(Record{"name":"Bob"}).Executor().Exec() func (id *InsertDataset) Executor() (executor exec.QueryExecutor) { if id.hooks != nil { - id.SetError(id.hooks.Before(id)) + if qf, err := id.hooks.UseTx(id); err != nil { + id.SetError(err) + } else if qf != nil { + id.queryFactory = qf + } + if herr := id.hooks.Before(id); herr != nil { + id.SetError(herr) + } } executor = id.queryFactory.FromSQLBuilder(id.insertSQLBuilder()) if id.hooks != nil { executor.Hook(func(result interface{}) error { - return id.hooks.After(id, result) + if herr := id.hooks.After(id, result); herr != nil { + return herr + } + return nil }) } return executor diff --git a/select_dataset.go b/select_dataset.go index 35a829f..b89de85 100644 --- a/select_dataset.go +++ b/select_dataset.go @@ -573,12 +573,17 @@ func (sd *SelectDataset) ToSQL() (sql string, params []interface{}, err error) { // See Dataset#ToUpdateSQL for arguments func (sd *SelectDataset) Executor() (executor exec.QueryExecutor) { if sd.hooks != nil { - sd.SetError(sd.hooks.Before(sd)) + if herr := sd.hooks.Before(sd); herr != nil { + sd.SetError(herr) + } } executor = sd.queryFactory.FromSQLBuilder(sd.selectSQLBuilder()) if sd.hooks != nil { executor.Hook(func(result interface{}) error { - return sd.hooks.After(sd, result) + if herr := sd.hooks.After(sd, result); herr != nil { + return herr + } + return nil }) } return executor diff --git a/update_dataset.go b/update_dataset.go index c7636e5..e16f6bb 100644 --- a/update_dataset.go +++ b/update_dataset.go @@ -236,12 +236,22 @@ func (ud *UpdateDataset) ReturnsColumns() bool { // db.Update("test").Set(Record{"name":"Bob", update: time.Now()}).Executor() func (ud *UpdateDataset) Executor() (executor exec.QueryExecutor) { if ud.hooks != nil { - ud.SetError(ud.hooks.Before(ud)) + if qf, err := ud.hooks.UseTx(ud); err != nil { + ud.SetError(err) + } else if qf != nil { + ud.queryFactory = qf + } + if herr := ud.hooks.Before(ud); herr != nil { + ud.SetError(herr) + } } executor = ud.queryFactory.FromSQLBuilder(ud.updateSQLBuilder()) if ud.hooks != nil { executor.Hook(func(result interface{}) error { - return ud.hooks.After(ud, result) + if herr := ud.hooks.After(ud, result); herr != nil { + return herr + } + return nil }) } return executor