From 7bc041e90808ef46f286225ab781c7545ce9c55e Mon Sep 17 00:00:00 2001 From: chenquan Date: Mon, 19 Sep 2022 22:14:09 +0800 Subject: [PATCH] feat: change interface --- conn.go | 12 ++++++++++++ conn_test.go | 7 +++++++ hook.go | 2 ++ hook_test.go | 10 ++++++++++ multi_hook.go | 16 ++++++++++++++++ 5 files changed, 47 insertions(+) diff --git a/conn.go b/conn.go index 88388e6..952d2eb 100644 --- a/conn.go +++ b/conn.go @@ -18,6 +18,18 @@ type conn struct { ConnHook } +func (c *conn) Close() (err error) { + ctx, err := c.BeforeClose(context.Background(), nil) + defer func() { + _, err = c.AfterClose(ctx, err) + }() + if err != nil { + return err + } + + return c.Conn.Close() +} + func (c *conn) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (result driver.Result, err error) { ctx, query, args, err = c.BeforeExecContext(ctx, query, args, nil) defer func() { diff --git a/conn_test.go b/conn_test.go index f4f1d71..d34bf04 100644 --- a/conn_test.go +++ b/conn_test.go @@ -244,7 +244,14 @@ func Test_conn_QueryContext(t *testing.T) { assert.Contains(t, s, "BeforeQueryContext") assert.Contains(t, s, "AfterQueryContext") }) +} +func TestMultiHook_Close(t *testing.T) { + c, m := createMockConn() + err := c.Close() + assert.NoError(t, err) + assert.Contains(t, m.String(), "AfterClose") + assert.Contains(t, m.String(), "BeforeClose") } func Test_namedValueToValue(t *testing.T) { diff --git a/hook.go b/hook.go index 6718b3e..18cc6d4 100644 --- a/hook.go +++ b/hook.go @@ -24,6 +24,8 @@ type ( BeforePrepareContext(ctx context.Context, query string, err error) (context.Context, string, error) AfterPrepareContext(ctx context.Context, query string, ds driver.Stmt, err error) (context.Context, driver.Stmt, error) + BeforeClose(ctx context.Context, err error) (context.Context, error) + AfterClose(ctx context.Context, err error) (context.Context, error) } ConnectorHook interface { BeforeConnect(ctx context.Context, err error) (context.Context, error) diff --git a/hook_test.go b/hook_test.go index bd9e99d..be3788d 100644 --- a/hook_test.go +++ b/hook_test.go @@ -12,6 +12,16 @@ type mockHook struct { Args []string } +func (m *mockHook) BeforeClose(ctx context.Context, err error) (context.Context, error) { + m.Write("BeforeClose") + return ctx, err +} + +func (m *mockHook) AfterClose(ctx context.Context, err error) (context.Context, error) { + m.Write("AfterClose") + return ctx, err +} + func (m *mockHook) BeforeConnect(ctx context.Context, err error) (context.Context, error) { m.Write("BeforeConnect") return ctx, err diff --git a/multi_hook.go b/multi_hook.go index 831f95d..12aec08 100644 --- a/multi_hook.go +++ b/multi_hook.go @@ -93,6 +93,22 @@ func (h *multiHook) AfterPrepareContext(ctx context.Context, query string, s dri return ctx, s, err } +func (h *multiHook) BeforeClose(ctx context.Context, err error) (context.Context, error) { + for _, hook := range h.hooks { + ctx, err = hook.BeforeClose(ctx, err) + } + + return ctx, err +} + +func (h *multiHook) AfterClose(ctx context.Context, err error) (context.Context, error) { + for i := len(h.hooks) - 1; i >= 0; i-- { + ctx, err = h.hooks[i].AfterClose(ctx, err) + } + + return ctx, err +} + func (h *multiHook) BeforeCommit(ctx context.Context, err error) (context.Context, error) { for _, hook := range h.hooks { ctx, err = hook.BeforeCommit(ctx, err)