Golang中如何Mock pgx驱动,为调用DB函数的Handler编写单元测试?
问题描述
我正在尝试为Golang的HTTP Handler编写单元测试,该Handler会调用TicketService的CreateTicketOption方法,而该方法通过pgx驱动与数据库交互。执行测试时出现panic,报错根源是测试时没有初始化数据库连接,同时希望分别对Handler和业务逻辑单独编写单元测试,请教如何Mock这类代码。
相关代码
Handler代码
func (s *Server) handleCreateTicketOption(w http.ResponseWriter, r *http.Request) { var t ticket.Ticket body, err := ioutil.ReadAll(r.Body) if err != nil { http.Error(w, er.ErrInternal.Error(), http.StatusInternalServerError) return } err = json.Unmarshal(body, &t) if err != nil { http.Error(w, er.ErrInvalidData.Error(), http.StatusBadRequest) return } ticket, err := s.TicketService.CreateTicketOption(r.Context(), t) if err != nil { http.Error(w, er.ErrInternal.Error(), http.StatusInternalServerError) return } res, err := json.Marshal(ticket) if err != nil { http.Error(w, er.ErrInternal.Error(), http.StatusInternalServerError) return } log.Printf("%v tickets allocated with name %v\n", t.Allocation, t.Name) s.sendResponse(w, res, http.StatusOK) }
数据库交互逻辑代码
func (t *TicketService) CreateTicketOption(ctx context.Context, ticket ticket.Ticket) (*ticket.Ticket, error) { tx, err := t.db.dbPool.Begin(ctx) if err != nil { return nil, er.ErrInternal } defer tx.Rollback(ctx) var id int err = tx.QueryRow(ctx, `INSERT INTO ticket (name, description, allocation) VALUES ($1, $2, $3) RETURNING id`, ticket.Name, ticket.Description, ticket.Allocation).Scan(&id) if err != nil { return nil, er.ErrInternal } ticket.Id = id return &ticket, tx.Commit(ctx) }
我编写的Handler单元测试代码
func TestCreateTicketOptionHandler(t *testing.T) { caseExpected, _ := json.Marshal(&ticket.Ticket{Id: 1, Name: "baris", Description: "test-desc", Allocation: 10}) srv := NewServer() // expected := [][]byte{ // _, _ = json.Marshal(&ticket.Ticket{Id: 1, Name: "baris", Description: "test-desc", Allocation: 20}), // // json.Marshal(&ticket.Ticket{Id: 1, Name: "baris", Description: "test-desc", Allocation: 20}) // } tt := []struct { name string entry *ticket.Ticket want []byte code int }{ { "valid", &ticket.Ticket{Name: "baris", Description: "test-desc", Allocation: 10}, caseExpected, http.StatusOK, }, } var buf bytes.Buffer for _, tc := range tt { t.Run(tc.name, func(t *testing.T) { json.NewEncoder(&buf).Encode(tc.entry) req, err := http.NewRequest(http.MethodPost, "/ticket_options", &buf) log.Println("1") if err != nil { log.Println("2") t.Fatalf("could not create request: %v", err) } log.Println("3") rec := httptest.NewRecorder() log.Println("4") srv.handleCreateTicketOption(rec, req) log.Println("5") if rec.Code != tc.code { t.Fatalf("got status %d, want %v", rec.Code, tc.code) } log.Println("6") if reflect.DeepEqual(rec.Body.Bytes(), tc.want) { log.Println("7") t.Fatalf("NAME:%v, got %v, want %v", tc.name, rec.Body.Bytes(), tc.want) } }) } }
报错信息
github.com/bariis/gowit-case-study/psql.(*TicketService).CreateTicketOption(0xc000061348, {0x1485058, 0xc0000260c0}, {0x0, {0xc000026dd0, 0x5}, {0xc000026dd5, 0x9}, 0xa}) /Users/barisertas/workspace/gowit-case-study/psql/ticket.go:24 +0x125 github.com/bariis/gowit-case-study/http.(*Server).handleCreateTicketOption(0xc000061340, {0x1484bf0, 0xc000153280}, 0xc00018e000) /Users/barisertas/workspace/gowit-case-study/http/ticket.go:77 +0x10b github.com/bariis/gowit-case-study/http.TestCreateTicketOptionHandler.func2(0xc000119860) /Users/barisertas/workspace/gowit-case-study/http/ticket_test.go:80 +0x305
报错位置:
psql/ticket.go:24:tx, err := t.db.dbPool.Begin(ctx)http/ticket.go:77:ticket, err := s.TicketService.CreateTicketOption(r.Context(), t)http/ticket_test.go:80:srv.handleCreateTicketOption(rec, req)
解决方案
要分别测试Handler和业务逻辑,核心是通过接口抽象解耦依赖,再配合Mock实现隔离测试。
一、Handler单元测试:Mock TicketService
Handler依赖TicketService,我们先给TicketService定义一个接口,让Server依赖这个接口而非具体实现,测试时就能替换成Mock对象。
步骤1:定义TicketService接口
在Handler所在包(比如http包)中定义接口:
// 只保留Handler需要用到的方法 type TicketService interface { CreateTicketOption(ctx context.Context, ticket ticket.Ticket) (*ticket.Ticket, error) } // 修改Server结构体,依赖接口而非具体实现 type Server struct { TicketService TicketService // 其他原有字段... }
步骤2:创建Mock TicketService
推荐用testify/mock库快速生成Mock,也可以手动实现:
方式1:使用testify/mock
先安装依赖:go get github.com/stretchr/testify/mock
创建Mock结构体:
import "github.com/stretchr/testify/mock" type MockTicketService struct { mock.Mock } func (m *MockTicketService) CreateTicketOption(ctx context.Context, ticket ticket.Ticket) (*ticket.Ticket, error) { args := m.Called(ctx, ticket) return args.Get(0).(*ticket.Ticket), args.Error(1) }
方式2:手动实现Mock
type MockTicketService struct { MockCreate func(ctx context.Context, ticket ticket.Ticket) (*ticket.Ticket, error) } func (m *MockTicketService) CreateTicketOption(ctx context.Context, ticket ticket.Ticket) (*ticket.Ticket, error) { return m.MockCreate(ctx, ticket) }
步骤3:编写Handler测试
修改测试代码,用Mock替换真实的TicketService:
func TestCreateTicketOptionHandler(t *testing.T) { caseExpected, _ := json.Marshal(&ticket.Ticket{Id: 1, Name: "baris", Description: "test-desc", Allocation: 10}) tt := []struct { name string entry *ticket.Ticket mockResp *ticket.Ticket mockErr error wantBody []byte wantCode int }{ { name: "valid request", entry: &ticket.Ticket{Name: "baris", Description: "test-desc", Allocation: 10}, mockResp: &ticket.Ticket{Id: 1, Name: "baris", Description: "test-desc", Allocation: 10}, mockErr: nil, wantBody: caseExpected, wantCode: http.StatusOK, }, { name: "service returns error", entry: &ticket.Ticket{Name: "baris", Description: "test-desc", Allocation: 10}, mockResp: nil, mockErr: er.ErrInternal, wantBody: []byte(er.ErrInternal.Error()), wantCode: http.StatusInternalServerError, }, } for _, tc := range tt { t.Run(tc.name, func(t *testing.T) { // 初始化Mock mockTS := new(MockTicketService) mockTS.On("CreateTicketOption", mock.Anything, *tc.entry).Return(tc.mockResp, tc.mockErr) // 创建Server并注入Mock srv := &Server{ TicketService: mockTS, } // 构造请求 var buf bytes.Buffer json.NewEncoder(&buf).Encode(tc.entry) req, err := http.NewRequest(http.MethodPost, "/ticket_options", &buf) if err != nil { t.Fatalf("could not create request: %v", err) } rec := httptest.NewRecorder() srv.handleCreateTicketOption(rec, req) // 验证状态码 if rec.Code != tc.wantCode { t.Fatalf("got status %d, want %d", rec.Code, tc.wantCode) } // 验证响应体(处理json.Marshal的换行差异) gotBody := bytes.TrimSpace(rec.Body.Bytes()) wantBody := bytes.TrimSpace(tc.wantBody) if !reflect.DeepEqual(gotBody, wantBody) { t.Fatalf("got body %q, want %q", gotBody, wantBody) } // 验证Mock方法是否被正确调用 mockTS.AssertExpectations(t) }) } }
二、业务逻辑单元测试:Mock数据库连接
业务逻辑(TicketService)依赖pgx的数据库连接池,我们通过接口抽象pgx的核心方法,Mock数据库交互。
步骤1:抽象数据库连接接口
在psql包中定义接口,覆盖CreateTicketOption用到的方法:
import ( "context" "github.com/jackc/pgx/v4" ) // DBInterface 抽象数据库连接的核心方法 type DBInterface interface { Begin(ctx context.Context) (pgx.Tx, error) } // 修改TicketService结构体,依赖接口 type TicketService struct { db DBInterface }
步骤2:Mock数据库连接和事务
Mock需要实现DBInterface以及pgx的Tx、Row接口:
import ( "context" "github.com/jackc/pgx/v4" "github.com/stretchr/testify/mock" ) // MockDB 实现DBInterface type MockDB struct { mock.Mock } func (m *MockDB) Begin(ctx context.Context) (pgx.Tx, error) { args := m.Called(ctx) return args.Get(0).(pgx.Tx), args.Error(1) } // MockTx 实现pgx.Tx接口的核心方法 type MockTx struct { mock.Mock } func (m *MockTx) QueryRow(ctx context.Context, sql string, args ...interface{}) pgx.Row { callArgs := m.Called(ctx, sql, args) return callArgs.Get(0).(pgx.Row) } func (m *MockTx) Commit(ctx context.Context) error { args := m.Called(ctx) return args.Error(0) } func (m *MockTx) Rollback(ctx context.Context) error { args := m.Called(ctx) return args.Error(0) } // MockRow 实现pgx.Row接口 type MockRow struct { mock.Mock } func (m *MockRow) Scan(dest ...interface{}) error { args := m.Called(dest...) return args.Error(0) }
步骤3:编写业务逻辑测试
func TestCreateTicketOption(t *testing.T) { tt := []struct { name string inputTicket ticket.Ticket mockBeginErr error mockScanErr error mockCommitErr error wantTicket *ticket.Ticket wantErr error }{ { name: "success create ticket", inputTicket: ticket.Ticket{Name: "test", Description: "test desc", Allocation: 10}, mockBeginErr: nil, mockScanErr: nil, mockCommitErr: nil, wantTicket: &ticket.Ticket{Id: 1, Name: "test", Description: "test desc", Allocation: 10}, wantErr: nil, }, { name: "begin transaction failed", inputTicket: ticket.Ticket{Name: "test", Description: "test desc", Allocation: 10}, mockBeginErr: er.ErrInternal, wantTicket: nil, wantErr: er.ErrInternal, }, } for _, tc := range tt { t.Run(tc.name, func(t *testing.T) { // 初始化MockRow mockRow := new(MockRow) if tc.mockScanErr == nil { mockRow.On("Scan", mock.AnythingOfType("*int")).Run(func(args mock.Arguments) { idPtr := args.Get(0).(*int) *idPtr = 1 // 设置返回的ID }).Return(nil) } else { mockRow.On("Scan", mock.AnythingOfType("*int")).Return(tc.mockScanErr) } // 初始化MockTx mockTx := new(MockTx) mockTx.On("QueryRow", mock.Anything, `INSERT INTO ticket (name, description, allocation) VALUES ($1, $2, $3) RETURNING id`, []interface{}{tc.inputTicket.Name, tc.inputTicket.Description, tc.inputTicket.Allocation}).Return(mockRow) mockTx.On("Commit", mock.Anything).Return(tc.mockCommitErr) mockTx.On("Rollback", mock.Anything).Return(nil) // 初始化MockDB mockDB := new(MockDB) mockDB.On("Begin", mock.Anything).Return(mockTx, tc.mockBeginErr) // 创建TicketService并注入MockDB ts := &TicketService{db: mockDB} // 调用方法 result, err := ts.CreateTicketOption(context.Background(), tc.inputTicket) // 验证结果 if err != tc.wantErr { t.Fatalf("got err %v, want %v", err, tc.wantErr) } if !reflect.DeepEqual(result, tc.wantTicket) { t.Fatalf("got ticket %+v, want %+v", result, tc.wantTicket) } // 验证所有Mock调用 mockDB.AssertExpectations(t) mockTx.AssertExpectations(t) mockRow.AssertExpectations(t) }) } }
额外提示
- 测试时要注意
json.Marshal的输出可能带换行,用bytes.TrimSpace处理后再比较。 - 接口只保留需要的方法,避免过度抽象。
- 不想用第三方Mock库时,手动实现Mock结构体完全可行,适合简单场景。
内容的提问来源于stack exchange,提问作者tassador
相关产品推荐
相关产品推荐

