From 14f408e6c7b105c6a04408f5ca5ae1a388fee462 Mon Sep 17 00:00:00 2001 From: chenquan Date: Tue, 13 Sep 2022 17:09:53 +0800 Subject: [PATCH] test: add more tests --- README.md | 4 +- conn_test.go | 15 +++-- connector.go | 2 +- connector_test.go | 18 ++++++ driver.go | 16 ++--- driver_test.go | 44 +++++++++++-- hook.go | 158 +-------------------------------------------- hook_test.go | 106 +++++++++++++++++++++++++++---- multi_hook.go | 159 ++++++++++++++++++++++++++++++++++++++++++++++ stmt.go | 18 ++++-- tx.go | 16 +++-- 11 files changed, 350 insertions(+), 206 deletions(-) create mode 100644 connector_test.go create mode 100644 multi_hook.go diff --git a/README.md b/README.md index f814328..7060c46 100644 --- a/README.md +++ b/README.md @@ -9,9 +9,9 @@ go get -u github.com/chenquan/sqlplus ``` # usage + Implement the `sqlplus.Hook` interface and wrap it with `sqlplus.New(d driver.Driver, hook Hook) driver.Driver` # example -- [sqltrace](https://github.com/chenquan/sqltrace): A sql link tracking library, -suitable for any relational database such as MySQL, oracle, SQL Server, PostgreSQL,TiDB etc. \ No newline at end of file +- [sqltrace](https://github.com/chenquan/sqltrace): A sql tracing library, suitable for any relational database such as Sqlite3, MySQL, Oracle, SQL Server, PostgreSQL, TiDB, etc. \ No newline at end of file diff --git a/conn_test.go b/conn_test.go index 0113971..f4f1d71 100644 --- a/conn_test.go +++ b/conn_test.go @@ -13,7 +13,7 @@ var _ driver.Conn = (*mockConn)(nil) type mockConn struct { } -func (c *mockConn) Prepare(query string) (driver.Stmt, error) { +func (c *mockConn) Prepare(_ string) (driver.Stmt, error) { return &mockStmt{}, nil } @@ -33,7 +33,7 @@ type mockConnQueryer struct { driver.Conn } -func (m *mockConnQueryer) Query(query string, args []driver.Value) (driver.Rows, error) { +func (m *mockConnQueryer) Query(_ string, _ []driver.Value) (driver.Rows, error) { return nil, nil } @@ -45,7 +45,7 @@ type mockConnQueryerContext struct { driver.Conn } -func (m *mockConnExecer) Exec(query string, args []driver.Value) (driver.Result, error) { +func (m *mockConnExecer) Exec(_ string, _ []driver.Value) (driver.Result, error) { return nil, nil } @@ -57,7 +57,7 @@ type mockConnExecer struct { driver.Conn } -func (m *mockConnQueryerContext) QueryContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Rows, error) { +func (m *mockConnQueryerContext) QueryContext(_ context.Context, _ string, _ []driver.NamedValue) (driver.Rows, error) { return nil, nil } @@ -68,7 +68,7 @@ type mockConnExecerContext struct { driver.Conn } -func (m *mockConnExecerContext) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Result, error) { +func (m *mockConnExecerContext) ExecContext(_ context.Context, _ string, _ []driver.NamedValue) (driver.Result, error) { return nil, nil } @@ -79,7 +79,7 @@ type mockConnBeginTx struct { driver.Conn } -func (m *mockConnBeginTx) BeginTx(ctx context.Context, opts driver.TxOptions) (driver.Tx, error) { +func (m *mockConnBeginTx) BeginTx(_ context.Context, _ driver.TxOptions) (driver.Tx, error) { return &mockTx{}, nil } @@ -87,9 +87,10 @@ func (m *mockConnBeginTx) BeginTx(ctx context.Context, opts driver.TxOptions) (d func createMockConn() (*conn, *mockHook) { m := &mockHook{} + hooks := NewMultiHook(m) c := &conn{ Conn: &mockConn{}, - ConnHook: m, + ConnHook: hooks, } return c, m } diff --git a/connector.go b/connector.go index 172df63..338bdf9 100644 --- a/connector.go +++ b/connector.go @@ -30,5 +30,5 @@ func (c *connector) Connect(ctx context.Context) (dc driver.Conn, err error) { } func (c *connector) Driver() driver.Driver { - return &Driver{Driver: c.Connector.Driver(), Hook: c.ConnectorHook.(Hook)} + return &wrappedDriver{Driver: c.Connector.Driver(), Hook: c.ConnectorHook.(Hook)} } diff --git a/connector_test.go b/connector_test.go new file mode 100644 index 0000000..b36b55d --- /dev/null +++ b/connector_test.go @@ -0,0 +1,18 @@ +package sqlplus + +import ( + "context" + "database/sql/driver" +) + +var _ driver.Connector = (*mockConnector)(nil) + +type mockConnector struct{} + +func (m *mockConnector) Connect(_ context.Context) (driver.Conn, error) { + return &mockConn{}, nil +} + +func (m *mockConnector) Driver() driver.Driver { + return &mockDriver{} +} diff --git a/driver.go b/driver.go index c4fa650..65ddb03 100644 --- a/driver.go +++ b/driver.go @@ -5,29 +5,29 @@ import ( ) var ( - _ driver.Driver = (*Driver)(nil) - _ driver.DriverContext = (*DriverCtx)(nil) + _ driver.Driver = (*wrappedDriver)(nil) + _ driver.DriverContext = (*wrappedDriverCtx)(nil) ) -type DriverCtx struct { +type wrappedDriverCtx struct { driver.Driver Hook } -type Driver struct { +type wrappedDriver struct { driver.Driver Hook } func New(d driver.Driver, hook Hook) driver.Driver { if _, ok := d.(driver.DriverContext); ok { - return &DriverCtx{Driver: d, Hook: hook} + return &wrappedDriverCtx{Driver: d, Hook: hook} } - return &Driver{Driver: d, Hook: hook} + return &wrappedDriver{Driver: d, Hook: hook} } -func (d Driver) Open(name string) (driver.Conn, error) { +func (d *wrappedDriver) Open(name string) (driver.Conn, error) { c, err := d.Driver.Open(name) if err != nil { return nil, err @@ -38,7 +38,7 @@ func (d Driver) Open(name string) (driver.Conn, error) { // ----------------- -func (d DriverCtx) OpenConnector(name string) (driver.Connector, error) { +func (d *wrappedDriverCtx) OpenConnector(name string) (driver.Connector, error) { if dd, ok := d.Driver.(driver.DriverContext); ok { openConnector, err := dd.OpenConnector(name) if err != nil { diff --git a/driver_test.go b/driver_test.go index 45a837b..5376548 100644 --- a/driver_test.go +++ b/driver_test.go @@ -1,6 +1,7 @@ package sqlplus import ( + "context" "database/sql/driver" "testing" @@ -8,17 +9,48 @@ import ( ) var _ driver.Driver = (*mockDriver)(nil) +var _ driver.DriverContext = (*mockDriverCtx)(nil) +var _ driver.Driver = (*mockDriverCtx)(nil) -type mockDriver struct { +type ( + mockDriver struct{} + mockDriverCtx struct{} +) + +func (d *mockDriver) Open(_ string) (driver.Conn, error) { + return &mockConn{}, nil } -func (d *mockDriver) Open(name string) (driver.Conn, error) { +// ----------------- + +func (m *mockDriverCtx) OpenConnector(_ string) (driver.Connector, error) { + return &mockConnector{}, nil +} + +func (m *mockDriverCtx) Open(_ string) (driver.Conn, error) { return &mockConn{}, nil } func TestNew(t *testing.T) { - d := New(&mockDriver{}, &mockHook{}) - conn, err := d.Open("any") - assert.NoError(t, err) - assert.NotNil(t, conn) + t.Run("mockDriver", func(t *testing.T) { + d := New(&mockDriver{}, &mockHook{}) + conn, err := d.Open("any") + assert.NoError(t, err) + assert.NotNil(t, conn) + }) + + t.Run("mockDriverCtx", func(t *testing.T) { + d := New(&mockDriverCtx{}, NewMultiHook(&mockHook{})) + connect, err := d.Open("any") + assert.NoError(t, err) + assert.NotNil(t, connect) + + driverContext := d.(driver.DriverContext) + openConnector, err := driverContext.OpenConnector("any") + connect, err = openConnector.Connect(context.Background()) + assert.NoError(t, err) + assert.NotNil(t, connect) + assert.NotNil(t, openConnector.Driver()) + }) + } diff --git a/hook.go b/hook.go index bc530b6..6718b3e 100644 --- a/hook.go +++ b/hook.go @@ -14,16 +14,16 @@ type ( } ConnHook interface { BeforeExecContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) - AfterExecContext(ctx context.Context, query string, args []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) + AfterExecContext(ctx context.Context, query string, args []driver.NamedValue, dr driver.Result, err error) (context.Context, driver.Result, error) BeforeBeginTx(ctx context.Context, opts driver.TxOptions, err error) (context.Context, driver.TxOptions, error) - AfterBeginTx(ctx context.Context, opts driver.TxOptions, dd driver.Tx, err error) (context.Context, driver.Tx, error) + AfterBeginTx(ctx context.Context, opts driver.TxOptions, dt driver.Tx, err error) (context.Context, driver.Tx, error) BeforeQueryContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) AfterQueryContext(ctx context.Context, query string, args []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) BeforePrepareContext(ctx context.Context, query string, err error) (context.Context, string, error) - AfterPrepareContext(ctx context.Context, query string, s driver.Stmt, err error) (context.Context, driver.Stmt, error) + AfterPrepareContext(ctx context.Context, query string, ds driver.Stmt, err error) (context.Context, driver.Stmt, error) } ConnectorHook interface { BeforeConnect(ctx context.Context, err error) (context.Context, error) @@ -41,156 +41,4 @@ type ( BeforeStmtExecContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, []driver.NamedValue, error) AfterStmtExecContext(ctx context.Context, query string, args []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) } - Hooks struct { - hooks []Hook - } ) - -func NewMultiHook(hooks ...Hook) Hook { - return &Hooks{hooks: hooks} -} - -func (h *Hooks) BeforeConnect(ctx context.Context, err error) (context.Context, error) { - for _, hook := range h.hooks { - ctx, err = hook.BeforeConnect(ctx, err) - } - - return ctx, err -} - -func (h *Hooks) AfterConnect(ctx context.Context, dc driver.Conn, err error) (context.Context, driver.Conn, error) { - for _, hook := range h.hooks { - ctx, dc, err = hook.AfterConnect(ctx, dc, err) - } - - return ctx, dc, err -} - -func (h *Hooks) BeforeExecContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) { - for _, hook := range h.hooks { - ctx, query, args, err = hook.BeforeExecContext(ctx, query, args, err) - } - - return ctx, query, args, err -} - -func (h *Hooks) AfterExecContext(ctx context.Context, query string, args []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) { - for _, hook := range h.hooks { - ctx, r, err = hook.AfterExecContext(ctx, query, args, r, err) - } - - return ctx, r, err -} - -func (h *Hooks) BeforeBeginTx(ctx context.Context, opts driver.TxOptions, err error) (context.Context, driver.TxOptions, error) { - for _, hook := range h.hooks { - ctx, opts, err = hook.BeforeBeginTx(ctx, opts, err) - } - - return ctx, opts, err -} - -func (h *Hooks) AfterBeginTx(ctx context.Context, opts driver.TxOptions, dd driver.Tx, err error) (context.Context, driver.Tx, error) { - for _, hook := range h.hooks { - ctx, dd, err = hook.AfterBeginTx(ctx, opts, dd, err) - } - - return ctx, dd, err -} - -func (h *Hooks) BeforeQueryContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) { - for _, hook := range h.hooks { - ctx, query, args, err = hook.BeforeQueryContext(ctx, query, args, err) - } - - return ctx, query, args, err -} - -func (h *Hooks) AfterQueryContext(ctx context.Context, query string, args []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) { - for _, hook := range h.hooks { - ctx, rows, err = hook.AfterQueryContext(ctx, query, args, rows, err) - - } - - return ctx, rows, err -} - -func (h *Hooks) BeforePrepareContext(ctx context.Context, query string, err error) (context.Context, string, error) { - for _, hook := range h.hooks { - ctx, query, err = hook.BeforePrepareContext(ctx, query, err) - } - - return ctx, query, err -} - -func (h *Hooks) AfterPrepareContext(ctx context.Context, query string, s driver.Stmt, err error) (context.Context, driver.Stmt, error) { - for _, hook := range h.hooks { - ctx, s, err = hook.AfterPrepareContext(ctx, query, s, err) - } - - return ctx, s, err -} - -func (h *Hooks) BeforeCommit(ctx context.Context, err error) (context.Context, error) { - for _, hook := range h.hooks { - ctx, err = hook.BeforeCommit(ctx, err) - } - - return ctx, err -} - -func (h *Hooks) AfterCommit(ctx context.Context, err error) (context.Context, error) { - for _, hook := range h.hooks { - ctx, err = hook.AfterCommit(ctx, err) - } - - return ctx, err -} - -func (h *Hooks) BeforeRollback(ctx context.Context, err error) (context.Context, error) { - for _, hook := range h.hooks { - ctx, err = hook.BeforeRollback(ctx, err) - } - - return ctx, err -} - -func (h *Hooks) AfterRollback(ctx context.Context, err error) (context.Context, error) { - for _, hook := range h.hooks { - ctx, err = hook.AfterRollback(ctx, err) - } - - return ctx, err -} - -func (h *Hooks) BeforeStmtQueryContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, []driver.NamedValue, error) { - for _, hook := range h.hooks { - ctx, args, err = hook.BeforeStmtQueryContext(ctx, query, args, err) - } - - return ctx, args, err -} - -func (h *Hooks) AfterStmtQueryContext(ctx context.Context, query string, args []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) { - for _, hook := range h.hooks { - ctx, rows, err = hook.AfterStmtQueryContext(ctx, query, args, rows, err) - } - - return ctx, rows, err -} - -func (h *Hooks) BeforeStmtExecContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, []driver.NamedValue, error) { - for _, hook := range h.hooks { - ctx, args, err = hook.BeforeStmtExecContext(ctx, query, args, err) - } - - return ctx, args, err -} - -func (h *Hooks) AfterStmtExecContext(ctx context.Context, query string, args []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) { - for _, hook := range h.hooks { - ctx, r, err = hook.AfterStmtExecContext(ctx, query, args, r, err) - } - - return ctx, r, err -} diff --git a/hook_test.go b/hook_test.go index e9d98a4..bd9e99d 100644 --- a/hook_test.go +++ b/hook_test.go @@ -27,7 +27,7 @@ func (m *mockHook) BeforeExecContext(ctx context.Context, query string, args []d return ctx, query, args, err } -func (m *mockHook) AfterExecContext(ctx context.Context, query string, args []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) { +func (m *mockHook) AfterExecContext(ctx context.Context, _ string, _ []driver.NamedValue, _ driver.Result, err error) (context.Context, driver.Result, error) { m.Write("AfterExecContext") return ctx, nil, err } @@ -37,10 +37,9 @@ func (m *mockHook) BeforeBeginTx(ctx context.Context, opts driver.TxOptions, err return ctx, opts, err } -func (m *mockHook) AfterBeginTx(ctx context.Context, opts driver.TxOptions, dd driver.Tx, err error) (context.Context, driver.Tx, error) { +func (m *mockHook) AfterBeginTx(ctx context.Context, _ driver.TxOptions, dt driver.Tx, err error) (context.Context, driver.Tx, error) { m.Write("AfterBeginTx") - return ctx, dd, err - + return ctx, dt, err } func (m *mockHook) BeforeQueryContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) { @@ -48,7 +47,7 @@ func (m *mockHook) BeforeQueryContext(ctx context.Context, query string, args [] return ctx, query, args, err } -func (m *mockHook) AfterQueryContext(ctx context.Context, query string, args []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) { +func (m *mockHook) AfterQueryContext(ctx context.Context, _ string, _ []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) { m.Write("AfterQueryContext") return ctx, rows, err } @@ -58,49 +57,128 @@ func (m *mockHook) BeforePrepareContext(ctx context.Context, query string, err e return ctx, query, err } -func (m *mockHook) AfterPrepareContext(ctx context.Context, query string, s driver.Stmt, err error) (context.Context, driver.Stmt, error) { +func (m *mockHook) AfterPrepareContext(ctx context.Context, _ string, ds driver.Stmt, err error) (context.Context, driver.Stmt, error) { m.Write("AfterPrepareContext") - return ctx, s, err + return ctx, ds, err } func (m *mockHook) BeforeCommit(ctx context.Context, err error) (context.Context, error) { m.Write("BeforeCommit") + txContext := TxContextFromContext(ctx) + if txContext == nil { + panic("txContext is nil") + } + + prepareContext := PrepareContextFromContext(ctx) + if prepareContext != nil { + panic("prepareContext is not nil") + } + return ctx, err } func (m *mockHook) AfterCommit(ctx context.Context, err error) (context.Context, error) { m.Write("AfterCommit") + txContext := TxContextFromContext(ctx) + if txContext == nil { + panic("txContext is nil") + } + + prepareContext := PrepareContextFromContext(ctx) + if prepareContext != nil { + panic("prepareContext is not nil") + } + return ctx, err } func (m *mockHook) BeforeRollback(ctx context.Context, err error) (context.Context, error) { m.Write("BeforeRollback") + txContext := TxContextFromContext(ctx) + if txContext == nil { + panic("txContext is nil") + } + + prepareContext := PrepareContextFromContext(ctx) + if prepareContext != nil { + panic("prepareContext is not nil") + } + return ctx, err } func (m *mockHook) AfterRollback(ctx context.Context, err error) (context.Context, error) { m.Write("AfterRollback") - return ctx, err + txContext := TxContextFromContext(ctx) + if txContext == nil { + panic("txContext is nil") + } + prepareContext := PrepareContextFromContext(ctx) + if prepareContext != nil { + panic("prepareContext is not nil") + } + + return ctx, err } -func (m *mockHook) BeforeStmtQueryContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) { +func (m *mockHook) BeforeStmtQueryContext(ctx context.Context, _ string, args []driver.NamedValue, err error) (context.Context, []driver.NamedValue, error) { m.Write("BeforeStmtQueryContext") - return ctx, query, args, err + prepareContext := PrepareContextFromContext(ctx) + if prepareContext == nil { + panic("prepareContext is nil") + } + + txContext := TxContextFromContext(ctx) + if txContext != nil { + panic("txContext is not nil") + } + + return ctx, args, err } -func (m *mockHook) AfterStmtQueryContext(ctx context.Context, query string, args []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) { +func (m *mockHook) AfterStmtQueryContext(ctx context.Context, _ string, _ []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) { m.Write("AfterStmtQueryContext") + prepareContext := PrepareContextFromContext(ctx) + if prepareContext == nil { + panic("prepareContext is nil") + } + + txContext := TxContextFromContext(ctx) + if txContext != nil { + panic("txContext is not nil") + } + return ctx, rows, err } -func (m *mockHook) BeforeStmtExecContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) { +func (m *mockHook) BeforeStmtExecContext(ctx context.Context, _ string, args []driver.NamedValue, err error) (context.Context, []driver.NamedValue, error) { m.Write("BeforeStmtExecContext") - return ctx, query, args, err + prepareContext := PrepareContextFromContext(ctx) + if prepareContext == nil { + panic("prepareContext is nil") + } + + txContext := TxContextFromContext(ctx) + if txContext != nil { + panic("txContext is not nil") + } + + return ctx, args, err } -func (m *mockHook) AfterStmtExecContext(ctx context.Context, query string, args []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) { +func (m *mockHook) AfterStmtExecContext(ctx context.Context, _ string, _ []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) { m.Write("AfterStmtExecContext") + prepareContext := PrepareContextFromContext(ctx) + if prepareContext == nil { + panic("prepareContext is nil") + } + + txContext := TxContextFromContext(ctx) + if txContext != nil { + panic("txContext is not nil") + } + return ctx, r, err } diff --git a/multi_hook.go b/multi_hook.go new file mode 100644 index 0000000..0525730 --- /dev/null +++ b/multi_hook.go @@ -0,0 +1,159 @@ +package sqlplus + +import ( + "context" + "database/sql/driver" +) + +type multiHook struct { + hooks []Hook +} + +func NewMultiHook(hooks ...Hook) Hook { + return &multiHook{hooks: hooks} +} + +func (h *multiHook) BeforeConnect(ctx context.Context, err error) (context.Context, error) { + for _, hook := range h.hooks { + ctx, err = hook.BeforeConnect(ctx, err) + } + + return ctx, err +} + +func (h *multiHook) AfterConnect(ctx context.Context, dc driver.Conn, err error) (context.Context, driver.Conn, error) { + for _, hook := range h.hooks { + ctx, dc, err = hook.AfterConnect(ctx, dc, err) + } + + return ctx, dc, err +} + +func (h *multiHook) BeforeExecContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) { + for _, hook := range h.hooks { + ctx, query, args, err = hook.BeforeExecContext(ctx, query, args, err) + } + + return ctx, query, args, err +} + +func (h *multiHook) AfterExecContext(ctx context.Context, query string, args []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) { + for _, hook := range h.hooks { + ctx, r, err = hook.AfterExecContext(ctx, query, args, r, err) + } + + return ctx, r, err +} + +func (h *multiHook) BeforeBeginTx(ctx context.Context, opts driver.TxOptions, err error) (context.Context, driver.TxOptions, error) { + for _, hook := range h.hooks { + ctx, opts, err = hook.BeforeBeginTx(ctx, opts, err) + } + + return ctx, opts, err +} + +func (h *multiHook) AfterBeginTx(ctx context.Context, opts driver.TxOptions, dd driver.Tx, err error) (context.Context, driver.Tx, error) { + for _, hook := range h.hooks { + ctx, dd, err = hook.AfterBeginTx(ctx, opts, dd, err) + } + + return ctx, dd, err +} + +func (h *multiHook) BeforeQueryContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, string, []driver.NamedValue, error) { + for _, hook := range h.hooks { + ctx, query, args, err = hook.BeforeQueryContext(ctx, query, args, err) + } + + return ctx, query, args, err +} + +func (h *multiHook) AfterQueryContext(ctx context.Context, query string, args []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) { + for _, hook := range h.hooks { + ctx, rows, err = hook.AfterQueryContext(ctx, query, args, rows, err) + + } + + return ctx, rows, err +} + +func (h *multiHook) BeforePrepareContext(ctx context.Context, query string, err error) (context.Context, string, error) { + for _, hook := range h.hooks { + ctx, query, err = hook.BeforePrepareContext(ctx, query, err) + } + + return ctx, query, err +} + +func (h *multiHook) AfterPrepareContext(ctx context.Context, query string, s driver.Stmt, err error) (context.Context, driver.Stmt, error) { + for _, hook := range h.hooks { + ctx, s, err = hook.AfterPrepareContext(ctx, query, s, err) + } + + return ctx, s, err +} + +func (h *multiHook) BeforeCommit(ctx context.Context, err error) (context.Context, error) { + for _, hook := range h.hooks { + ctx, err = hook.BeforeCommit(ctx, err) + } + + return ctx, err +} + +func (h *multiHook) AfterCommit(ctx context.Context, err error) (context.Context, error) { + for _, hook := range h.hooks { + ctx, err = hook.AfterCommit(ctx, err) + } + + return ctx, err +} + +func (h *multiHook) BeforeRollback(ctx context.Context, err error) (context.Context, error) { + for _, hook := range h.hooks { + ctx, err = hook.BeforeRollback(ctx, err) + } + + return ctx, err +} + +func (h *multiHook) AfterRollback(ctx context.Context, err error) (context.Context, error) { + for _, hook := range h.hooks { + ctx, err = hook.AfterRollback(ctx, err) + } + + return ctx, err +} + +func (h *multiHook) BeforeStmtQueryContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, []driver.NamedValue, error) { + for _, hook := range h.hooks { + ctx, args, err = hook.BeforeStmtQueryContext(ctx, query, args, err) + } + + return ctx, args, err +} + +func (h *multiHook) AfterStmtQueryContext(ctx context.Context, query string, args []driver.NamedValue, rows driver.Rows, err error) (context.Context, driver.Rows, error) { + for _, hook := range h.hooks { + ctx, rows, err = hook.AfterStmtQueryContext(ctx, query, args, rows, err) + } + + return ctx, rows, err +} + +func (h *multiHook) BeforeStmtExecContext(ctx context.Context, query string, args []driver.NamedValue, err error) (context.Context, []driver.NamedValue, error) { + for _, hook := range h.hooks { + ctx, args, err = hook.BeforeStmtExecContext(ctx, query, args, err) + } + + return ctx, args, err +} + +func (h *multiHook) AfterStmtExecContext(ctx context.Context, query string, args []driver.NamedValue, r driver.Result, err error) (context.Context, driver.Result, error) { + for _, hook := range h.hooks { + ctx, r, err = hook.AfterStmtExecContext(ctx, query, args, r, err) + } + + return ctx, r, err +} diff --git a/stmt.go b/stmt.go index b3a61fa..2c24919 100644 --- a/stmt.go +++ b/stmt.go @@ -11,7 +11,15 @@ var ( _ driver.StmtQueryContext = (*stmt)(nil) ) -type prepareContextKey struct{} +type ( + stmt struct { + driver.Stmt + query string + StmtHook + prepareContext context.Context + } + prepareContextKey struct{} +) func PrepareContextFromContext(ctx context.Context) context.Context { value := ctx.Value(prepareContextKey{}) @@ -22,12 +30,7 @@ func PrepareContextFromContext(ctx context.Context) context.Context { return nil } -type stmt struct { - driver.Stmt - query string - StmtHook - prepareContext context.Context -} +// ----------------- func (s *stmt) QueryContext(ctx context.Context, args []driver.NamedValue) (rows driver.Rows, err error) { query := s.query @@ -56,6 +59,7 @@ func (s *stmt) QueryContext(ctx context.Context, args []driver.NamedValue) (rows func (s *stmt) ExecContext(ctx context.Context, args []driver.NamedValue) (r driver.Result, err error) { query := s.query + ctx = context.WithValue(ctx, prepareContextKey{}, s.prepareContext) ctx, args, err = s.BeforeStmtExecContext(ctx, query, args, nil) defer func() { _, r, err = s.AfterStmtExecContext(ctx, query, args, r, err) diff --git a/tx.go b/tx.go index f6b94d8..ae6c91d 100644 --- a/tx.go +++ b/tx.go @@ -5,12 +5,14 @@ import ( "database/sql/driver" ) -type txContextKey struct{} -type tx struct { - driver.Tx - TxHook - txContext context.Context -} +type ( + tx struct { + driver.Tx + TxHook + txContext context.Context + } + txContextKey struct{} +) func TxContextFromContext(ctx context.Context) context.Context { value := ctx.Value(txContextKey{}) @@ -21,6 +23,8 @@ func TxContextFromContext(ctx context.Context) context.Context { return nil } +// ----------------- + func (t *tx) Commit() (err error) { ctx := context.WithValue(context.Background(), txContextKey{}, t.txContext) ctx, err = t.BeforeCommit(ctx, nil)