You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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)
        })
    }
}

额外提示

  1. 测试时要注意json.Marshal的输出可能带换行,用bytes.TrimSpace处理后再比较。
  2. 接口只保留需要的方法,避免过度抽象。
  3. 不想用第三方Mock库时,手动实现Mock结构体完全可行,适合简单场景。

内容的提问来源于stack exchange,提问作者tassador

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.20 02:50:33