Skip to content

Mock 与接口测试

单元测试的核心思想是「隔离」——只测当前被测对象的逻辑,不依赖数据库、网络、第三方服务等外部资源。要做到这点,最常见的手段就是 Mock(用替身替换真实依赖)。本篇将系统讲解 Go 中如何通过接口实现依赖注入、手动 Mock、gomock、mockery、testify/mock 等多种方式编写可测试的代码,并通过一个分层架构的完整示例演示 Repository 与 Service 层的测试套路。

一、依赖注入与可测试性

一段代码是否易于测试,很大程度上取决于它的依赖如何获取。来看两个反例:

go
// 反例 1:直接全局调用,无法替换
package service

import "database/sql"

func GetUser(id int) (User, error) {
	var u User
	err := db.QueryRow("SELECT ...", id).Scan(...)
	//         ↑ db 是某个全局变量
	return u, err
}
go
// 反例 2:在函数内部 new 依赖,无法替换
func NewOrderService() *OrderService {
	return &OrderService{
		repo: &MySQLRepo{dsn: "..."}, // 写死实现
	}
}

这两种写法的共同问题是:依赖是「写死」的,测试时无法替换。要测试它们,就必须真的连上一个 MySQL,这已经超出了单元测试的范畴。

依赖注入(Dependency Injection, DI)的核心做法是:依赖从外部传入,而不是内部创建

go
package service

// Repository 是一个接口,由外部传入实现。
type Repository interface {
	FindByID(id int) (User, error)
}

type UserService struct {
	repo Repository // 依赖接口,而不是具体实现
}

func NewUserService(repo Repository) *UserService {
	return &UserService{repo: repo}
}

func (s *UserService) GetUser(id int) (User, error) {
	return s.repo.FindByID(id)
}

这样测试时,可以传入一个 Mock 实现:

go
type fakeRepo struct {
	user User
	err  error
}

func (f *fakeRepo) FindByID(id int) (User, error) {
	return f.user, f.err
}

func TestUserService_GetUser(t *testing.T) {
	svc := NewUserService(&fakeRepo{user: User{ID: 1, Name: "Alice"}})
	got, _ := svc.GetUser(1)
	if got.Name != "Alice" {
		t.Errorf("got %q", got.Name)
	}
}

可测试性的关键:依赖接口,而不是依赖具体实现。这样 Mock 才有空间。

二、接口驱动的测试

Go 的接口是隐式实现(duck typing)——一个类型只要实现了接口里所有方法,就被认为实现了该接口,无需 implements 声明。这个特性让 Mock 变得极其自然:

  1. 在被测代码中定义依赖接口。
  2. 在生产代码中实现该接口(如 MySQLRepo)。
  3. 在测试代码中写一个 Mock 实现,也实现该接口。

来看一个完整的 Repository 模式示例:

go
// user/repository.go
package user

import "context"

// User 是业务对象。
type User struct {
	ID    int
	Name  string
	Email string
}

// Repository 是数据访问接口。
type Repository interface {
	FindByID(ctx context.Context, id int) (User, error)
	Create(ctx context.Context, u User) (User, error)
	Update(ctx context.Context, u User) error
	Delete(ctx context.Context, id int) error
}

生产实现:

go
// user/mysql_repo.go
package user

import (
	"context"
	"database/sql"
)

type MySQLRepo struct {
	db *sql.DB
}

func NewMySQLRepo(db *sql.DB) *MySQLRepo {
	return &MySQLRepo{db: db}
}

func (r *MySQLRepo) FindByID(ctx context.Context, id int) (User, error) {
	var u User
	err := r.db.QueryRowContext(ctx,
		"SELECT id, name, email FROM users WHERE id = ?", id,
	).Scan(&u.ID, &u.Name, &u.Email)
	return u, err
}

func (r *MySQLRepo) Create(ctx context.Context, u User) (User, error) {
	res, err := r.db.ExecContext(ctx,
		"INSERT INTO users (name, email) VALUES (?, ?)", u.Name, u.Email,
	)
	if err != nil {
		return User{}, err
	}
	id, _ := res.LastInsertId()
	u.ID = int(id)
	return u, nil
}

func (r *MySQLRepo) Update(ctx context.Context, u User) error {
	_, err := r.db.ExecContext(ctx,
		"UPDATE users SET name = ?, email = ? WHERE id = ?", u.Name, u.Email, u.ID,
	)
	return err
}

func (r *MySQLRepo) Delete(ctx context.Context, id int) error {
	_, err := r.db.ExecContext(ctx, "DELETE FROM users WHERE id = ?", id)
	return err
}

Service 层依赖 Repository 接口:

go
// user/service.go
package user

import (
	"context"
	"errors"
)

var ErrUserNotFound = errors.New("user not found")

type Service struct {
	repo Repository
}

func NewService(repo Repository) *Service {
	return &Service{repo: repo}
}

func (s *Service) Get(ctx context.Context, id int) (User, error) {
	u, err := s.repo.FindByID(ctx, id)
	if err != nil {
		return User{}, err
	}
	return u, nil
}

func (s *Service) Register(ctx context.Context, name, email string) (User, error) {
	if name == "" || email == "" {
		return User{}, errors.New("name and email are required")
	}
	return s.repo.Create(ctx, User{Name: name, Email: email})
}

func (s *Service) ChangeEmail(ctx context.Context, id int, newEmail string) error {
	u, err := s.repo.FindByID(ctx, id)
	if err != nil {
		return ErrUserNotFound
	}
	u.Email = newEmail
	return s.repo.Update(ctx, u)
}

测试时只需 Mock Repository 接口,无需连数据库。

三、手动 Mock 实现

最简单的 Mock 就是手写一个实现接口的 struct:

go
package user_test

import (
	"context"
	"errors"
	"testing"

	"example.com/user"
)

// fakeRepo 是手写的 Mock 实现。
type fakeRepo struct {
	users map[int]user.User
	err   error // 注入的错误,用于测试错误路径
}

func newFakeRepo() *fakeRepo {
	return &fakeRepo{users: make(map[int]user.User)}
}

func (r *fakeRepo) FindByID(ctx context.Context, id int) (user.User, error) {
	if r.err != nil {
		return user.User{}, r.err
	}
	u, ok := r.users[id]
	if !ok {
		return user.User{}, errors.New("not found")
	}
	return u, nil
}

func (r *fakeRepo) Create(ctx context.Context, u user.User) (user.User, error) {
	if r.err != nil {
		return user.User{}, r.err
	}
	u.ID = len(r.users) + 1
	r.users[u.ID] = u
	return u, nil
}

func (r *fakeRepo) Update(ctx context.Context, u user.User) error {
	if r.err != nil {
		return r.err
	}
	if _, ok := r.users[u.ID]; !ok {
		return errors.New("not found")
	}
	r.users[u.ID] = u
	return nil
}

func (r *fakeRepo) Delete(ctx context.Context, id int) error {
	if r.err != nil {
		return r.err
	}
	delete(r.users, id)
	return nil
}

func TestService_Get(t *testing.T) {
	repo := newFakeRepo()
	repo.users[1] = user.User{ID: 1, Name: "Alice", Email: "alice@example.com"}

	svc := user.NewService(repo)
	got, err := svc.Get(context.Background(), 1)
	if err != nil {
		t.Fatalf("unexpected error: %v", err)
	}
	if got.Name != "Alice" {
		t.Errorf("got %q, want Alice", got.Name)
	}
}

func TestService_GetNotFound(t *testing.T) {
	repo := newFakeRepo()
	svc := user.NewService(repo)

	_, err := svc.Get(context.Background(), 999)
	if err == nil {
		t.Fatal("expected error, got nil")
	}
}

func TestService_Register(t *testing.T) {
	repo := newFakeRepo()
	svc := user.NewService(repo)

	u, err := svc.Register(context.Background(), "Bob", "bob@example.com")
	if err != nil {
		t.Fatalf("unexpected: %v", err)
	}
	if u.ID != 1 {
		t.Errorf("ID = %d, want 1", u.ID)
	}
}

func TestService_ChangeEmail(t *testing.T) {
	repo := newFakeRepo()
	repo.users[1] = user.User{ID: 1, Name: "Alice", Email: "old@example.com"}
	svc := user.NewService(repo)

	err := svc.ChangeEmail(context.Background(), 1, "new@example.com")
	if err != nil {
		t.Fatalf("unexpected: %v", err)
	}
	if got := repo.users[1].Email; got != "new@example.com" {
		t.Errorf("email = %q, want new@example.com", got)
	}
}

手动 Mock 的优点:

  • 不需要任何外部工具,纯 Go 代码。
  • 实现逻辑完全可控,可以模拟「真实数据库行为」(如自增 ID)。
  • 复用性强:一份 Mock 可以被多个测试共享。

缺点:

  • 接口方法多时,写 Mock 实现很冗长。
  • 难以精确断言「方法被调用了几次、参数是什么」。

四、gomock 简介

gomock 是 Google 开源的 Mock 生成工具,由两部分组成:

  1. mockgen:命令行工具,根据接口定义自动生成 Mock 代码。
  2. gomock 包:运行时库,提供 EXPECT() API 编写期望。

安装:

bash
go install github.com/golang/mock/mockgen@latest

注:原 github.com/golang/mock 已迁移到 go.uber.org/mock,新项目推荐使用后者,API 完全兼容。

定义接口后,用 mockgen 生成:

bash
# reflect 模式(推荐):从源码反射生成
mockgen -source=user/repository.go -destination=user/mock_repository.go -package=user

# source 模式:指定包与接口名
mockgen example.com/user Repository > mock_repository.go

生成的 Mock 长这样(简化):

go
// Code generated by MockGen. DO NOT EDIT.
package user

import (
	"context"
	"reflect"

	"go.uber.org/mock/gomock"
)

type MockRepository struct {
	ctrl     *gomock.Controller
	recorder *MockRepositoryMockRecorder
}

func NewMockRepository(ctrl *gomock.Controller) *MockRepository {
	mock := &MockRepository{ctrl: ctrl}
	mock.recorder = &MockRepositoryMockRecorder{mock}
	return mock
}

func (m *MockRepository) EXPECT() *MockRepositoryMockRecorder {
	return m.recorder
}

func (m *MockRepository) FindByID(ctx context.Context, id int) (User, error) {
	m.ctrl.T.Helper()
	ret := m.ctrl.Call(m, "FindByID", ctx, id)
	// ...
	return ret[0].(User), ret[1].(error)
}

type MockRepositoryMockRecorder struct {
	mock *MockRepository
}

func (mr *MockRepositoryMockRecorder) FindByID(ctx context.Context, id int) *gomock.Call {
	mr.mock.ctrl.T.Helper()
	return mr.mock.ctrl.RecordCall(mr.mock.T, "FindByID", ctx, id)
}

使用 gomock 写测试:

go
package user_test

import (
	"context"
	"testing"

	"go.uber.org/mock/gomock"

	"example.com/user"
	// 这个 import 路径取决于 mockgen 生成的文件位置
	// 假设生成的文件位于 user 包下
)

func TestService_Get_gomock(t *testing.T) {
	ctrl := gomock.NewController(t)
	defer ctrl.Finish() // 校验所有期望都被满足

	mockRepo := user.NewMockRepository(ctrl)
	mockRepo.EXPECT().
		FindByID(gomock.Any(), 1).
		Return(user.User{ID: 1, Name: "Alice"}, nil).
		Times(1)

	svc := user.NewService(mockRepo)
	got, err := svc.Get(context.Background(), 1)
	if err != nil {
		t.Fatalf("unexpected: %v", err)
	}
	if got.Name != "Alice" {
		t.Errorf("got %q", got.Name)
	}
	// 调用结束时,ctrl.Finish 会校验 FindByID 被调用了恰好 1 次
}

func TestService_ChangeEmail_gomock(t *testing.T) {
	ctrl := gomock.NewController(t)
	defer ctrl.Finish()

	mockRepo := user.NewMockRepository(ctrl)
	//gomock 期望:先 FindByID,再 Update,按顺序
	gomock.InOrder(
		mockRepo.EXPECT().
			FindByID(gomock.Any(), 1).
			Return(user.User{ID: 1, Name: "Alice", Email: "old@example.com"}, nil),
		mockRepo.EXPECT().
			Update(gomock.Any(), user.User{ID: 1, Name: "Alice", Email: "new@example.com"}).
			Return(nil),
	)

	svc := user.NewService(mockRepo)
	err := svc.ChangeEmail(context.Background(), 1, "new@example.com")
	if err != nil {
		t.Fatalf("unexpected: %v", err)
	}
}

gomock 的核心 API:

API用途
mock.EXPECT().Method(args)设置方法期望
.Return(values...)设置返回值
.Times(n)期望被调用恰好 n 次
.AnyTimes()期望被调用任意次(含 0 次)
.MinTimes(n) / .MaxTimes(n)期望调用次数上下界
gomock.Any()参数匹配:任意值
gomock.Eq(x)参数匹配:等于 x
gomock.InOrder(calls...)设置多个调用的顺序
gomock.NewController(t)创建 Controller,Finish 校验期望
.Do(func) / .DoAndReturn(func)调用时执行回调(适合动态行为)

gomock 的优点:

  • 自动生成,省力。
  • 严格校验调用次数与顺序,能发现「方法没被调用」的问题。
  • 社区广泛使用,文档丰富。

缺点:

  • 需要维护生成命令与生成文件。
  • 过度严格的断言(Times(1))会让测试脆弱——稍微重构就被打破。

五、mockery 自动生成 Mock

mockery 是另一个流行的 Mock 生成工具,相比 mockgen 的特点:

  • 配置驱动(.mockery.yaml),不需要每次写命令行。
  • 生成的代码风格更现代。
  • 支持泛型接口(Go 1.18+)。

安装:

bash
go install github.com/vektra/mockery/v2@latest

在项目根目录创建 .mockery.yaml

yaml
with-expecter: true
filename: "mock_{{.InterfaceName | snakecase}}.go"
dir: "mocks/{{.PackageName}}"
mockname: "Mock{{.InterfaceName}}"
outpkg: "{{.PackageName}}"
packages:
  example.com/user:
    interfaces:
      Repository:

执行 mockery 命令即可生成。生成的 Mock 与 testify/mock 兼容,下面会介绍 testify/mock 的用法。

六、testify/mock 使用

testify/mock 是 testify 生态自带的 Mock 框架,无需生成代码(也可以生成),通过嵌入 mock.Mock 字段实现接口。

go
package user_test

import (
	"context"
	"testing"

	"github.com/stretchr/testify/assert"
	"github.com/stretchr/testify/mock"
	"github.com/stretchr/testify/require"

	"example.com/user"
)

// MockRepository 是 testify/mock 风格的 Mock。
type MockRepository struct {
	mock.Mock
}

func (m *MockRepository) FindByID(ctx context.Context, id int) (user.User, error) {
	args := m.Called(ctx, id)
	return args.Get(0).(user.User), args.Error(1)
}

func (m *MockRepository) Create(ctx context.Context, u user.User) (user.User, error) {
	args := m.Called(ctx, u)
	return args.Get(0).(user.User), args.Error(1)
}

func (m *MockRepository) Update(ctx context.Context, u user.User) error {
	args := m.Called(ctx, u)
	return args.Error(0)
}

func (m *MockRepository) Delete(ctx context.Context, id int) error {
	args := m.Called(ctx, id)
	return args.Error(0)
}

func TestService_Get_testify(t *testing.T) {
	mockRepo := new(MockRepository)
	mockRepo.On("FindByID", mock.Anything, 1).
		Return(user.User{ID: 1, Name: "Alice"}, nil)

	svc := user.NewService(mockRepo)
	got, err := svc.Get(context.Background(), 1)

	require.NoError(t, err)
	assert.Equal(t, "Alice", got.Name)

	mockRepo.AssertExpectations(t) // 校验所有 On 注册的期望都被调用
}

func TestService_ChangeEmail_testify(t *testing.T) {
	mockRepo := new(MockRepository)
	mockRepo.On("FindByID", mock.Anything, 1).
		Return(user.User{ID: 1, Name: "Alice", Email: "old@example.com"}, nil)
	mockRepo.On("Update", mock.Anything, user.User{ID: 1, Name: "Alice", Email: "new@example.com"}).
		Return(nil)

	svc := user.NewService(mockRepo)
	err := svc.ChangeEmail(context.Background(), 1, "new@example.com")

	require.NoError(t, err)
	// 也可以断言「FindByID 被调用 1 次」
	mockRepo.AssertNumberOfCalls(t, "FindByID", 1)
	mockRepo.AssertCalled(t, "Update", mock.Anything, user.User{ID: 1, Name: "Alice", Email: "new@example.com"})
}

func TestService_Get_Error(t *testing.T) {
	mockRepo := new(MockRepository)
	mockRepo.On("FindByID", mock.Anything, 999).
		Return(user.User{}, errors.New("db error"))

	svc := user.NewService(mockRepo)
	_, err := svc.Get(context.Background(), 999)

	require.Error(t, err)
	assert.Contains(t, err.Error(), "db error")
}

testify/mock 的核心 API:

API用途
m.On("Method", args...).Return(values...)注册一次期望
.Once() / .Twice() / .Times(n)期望调用次数
.Maybe()可选调用,未发生也不报错
mock.Anything匹配任意参数
m.Called(args...)在 Mock 实现中调用,触发期望
m.AssertExpectations(t)校验所有 On 都被实际调用
m.AssertCalled(t, "Method", args...)断言方法被以指定参数调用过
m.AssertNumberOfCalls(t, "Method", n)断言调用次数
m.AssertNotCalled(t, "Method")断言方法从未被调用

参数匹配:默认是「深度相等」。如需自定义匹配,可以实现 mock.ArgumentMatcher

go
mockRepo.On("FindByID", mock.Anything, mock.MatchedBy(func(id int) bool {
	return id > 0
})).Return(user.User{ID: 1, Name: "Alice"}, nil)

七、Mock 数据库层

数据库是单元测试里最常见的「外部依赖」。Mock 数据库层有几种思路:

方案 1:Mock Repository 接口

最常见。把 database/sql.DB 的访问封装在 Repository 接口后,Service 层只依赖接口。如上文所示。

方案 2:用 sqlmock 模拟 sql.DB

如果直接使用 *sql.DB,难以 Mock,因为 sql.DB 是结构体而非接口。github.com/DATA-DOG/go-sqlmock 库可以在不修改代码的前提下,模拟出 SQL 执行的结果:

go
package user_test

import (
	"context"
	"database/sql"
	"testing"

	"github.com/DATA-DOG/go-sqlmock"
	"github.com/stretchr/testify/assert"
	"github.com/stretchr/testify/require"

	"example.com/user"
)

func TestMySQLRepo_FindByID(t *testing.T) {
	db, mock, err := sqlmock.New()
	require.NoError(t, err)
	defer db.Close()

	repo := user.NewMySQLRepo(db)

	// 期望:执行 SELECT id, name, email FROM users WHERE id = ?,参数为 1
	rows := sqlmock.NewRows([]string{"id", "name", "email"}).
		AddRow(1, "Alice", "alice@example.com")
	mock.ExpectQuery("SELECT id, name, email FROM users WHERE id = ?").
		WithArgs(1).
		WillReturnRows(rows)

	got, err := repo.FindByID(context.Background(), 1)
	require.NoError(t, err)
	assert.Equal(t, "Alice", got.Name)

	// 校验所有期望都被满足
	require.NoError(t, mock.ExpectationsWereMet())
}

func TestMySQLRepo_Create(t *testing.T) {
	db, mock, err := sqlmock.New()
	require.NoError(t, err)
	defer db.Close()

	repo := user.NewMySQLRepo(db)

	mock.ExpectExec("INSERT INTO users").
		WithArgs("Bob", "bob@example.com").
		WillReturnResult(sqlmock.NewResult(42, 1))

	got, err := repo.Create(context.Background(), user.User{Name: "Bob", Email: "bob@example.com"})
	require.NoError(t, err)
	assert.Equal(t, 42, got.ID)
	require.NoError(t, mock.ExpectationsWereMet())
}

sqlmock 的优点是能验证 SQL 语句与参数,但缺点也很明显:测试与 SQL 文本强耦合,稍微改 SQL(比如换行、空格)就崩。

方案 3:用 SQLite 内存库做集成测试

更接近真实数据库的行为,但本质上是集成测试,下篇会专门讲。

八、Mock HTTP 客户端

调用外部 HTTP API 的代码也要 Mock。常见做法是把 HTTP 调用抽象成接口:

go
package weather

import "context"

// WeatherClient 是天气 API 客户端接口。
type WeatherClient interface {
	GetWeather(ctx context.Context, city string) (string, error)
}

type Service struct {
	client WeatherClient
}

func NewService(c WeatherClient) *Service {
	return &Service{client: c}
}

func (s *Service) GetWeatherReport(ctx context.Context, city string) (string, error) {
	w, err := s.client.GetWeather(ctx, city)
	if err != nil {
		return "", err
	}
	return "Today's weather: " + w, nil
}

测试用 Mock 实现:

go
package weather_test

import (
	"context"
	"errors"
	"testing"

	"github.com/stretchr/testify/assert"
	"github.com/stretchr/testify/require"

	"example.com/weather"
)

type fakeWeatherClient struct {
	weather string
	err     error
	called  bool
}

func (f *fakeWeatherClient) GetWeather(ctx context.Context, city string) (string, error) {
	f.called = true
	return f.weather, f.err
}

func TestService_GetWeatherReport(t *testing.T) {
	cases := []struct {
		name    string
		weather string
		err     error
		want    string
		wantErr bool
	}{
		{"晴天", "sunny, 25C", nil, "Today's weather: sunny, 25C", false},
		{"雨天", "rainy, 18C", nil, "Today's weather: rainy, 18C", false},
		{"错误", "", errors.New("api timeout"), "", true},
	}

	for _, c := range cases {
		t.Run(c.name, func(t *testing.T) {
			fake := &fakeWeatherClient{weather: c.weather, err: c.err}
			svc := weather.NewService(fake)

			got, err := svc.GetWeatherReport(context.Background(), "Beijing")
			if c.wantErr {
				require.Error(t, err)
				return
			}
			require.NoError(t, err)
			assert.Equal(t, c.want, got)
			assert.True(t, fake.called)
		})
	}
}

如果直接使用了 http.Client,也可以用 httptest.NewServer 来 Mock(下一篇 HTTP 测试详述)。

九、完整示例:分层架构的单元测试

把前面的概念整合,下面给出一个完整的「Repository + Service」分层示例,并附上每层的测试。

Repository 层 Mock(基于 sqlmock)

product.go

go
package product

import "context"

type Product struct {
	ID    int
	Name  string
	Price float64
}

type Repository interface {
	FindByID(ctx context.Context, id int) (Product, error)
	Create(ctx context.Context, p Product) (Product, error)
	List(ctx context.Context) ([]Product, error)
}

mysql_repo.go

go
package product

import (
	"context"
	"database/sql"
	"fmt"
)

type MySQLRepo struct {
	db *sql.DB
}

func NewMySQLRepo(db *sql.DB) *MySQLRepo {
	return &MySQLRepo{db: db}
}

func (r *MySQLRepo) FindByID(ctx context.Context, id int) (Product, error) {
	var p Product
	err := r.db.QueryRowContext(ctx,
		"SELECT id, name, price FROM products WHERE id = ?", id,
	).Scan(&p.ID, &p.Name, &p.Price)
	if err == sql.ErrNoRows {
		return Product{}, fmt.Errorf("product not found: id=%d", id)
	}
	return p, err
}

func (r *MySQLRepo) Create(ctx context.Context, p Product) (Product, error) {
	res, err := r.db.ExecContext(ctx,
		"INSERT INTO products (name, price) VALUES (?, ?)", p.Name, p.Price,
	)
	if err != nil {
		return Product{}, err
	}
	id, _ := res.LastInsertId()
	p.ID = int(id)
	return p, nil
}

func (r *MySQLRepo) List(ctx context.Context) ([]Product, error) {
	rows, err := r.db.QueryContext(ctx, "SELECT id, name, price FROM products")
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	var list []Product
	for rows.Next() {
		var p Product
		if err := rows.Scan(&p.ID, &p.Name, &p.Price); err != nil {
			return nil, err
		}
		list = append(list, p)
	}
	return list, rows.Err()
}

mysql_repo_test.go

go
package product_test

import (
	"context"
	"database/sql"
	"errors"
	"testing"

	"github.com/DATA-DOG/go-sqlmock"
	"github.com/stretchr/testify/assert"
	"github.com/stretchr/testify/require"

	"example.com/product"
)

func TestMySQLRepo_FindByID(t *testing.T) {
	db, mock, err := sqlmock.New()
	require.NoError(t, err)
	defer db.Close()

	repo := product.NewMySQLRepo(db)

	t.Run("found", func(t *testing.T) {
		rows := sqlmock.NewRows([]string{"id", "name", "price"}).
			AddRow(1, "Apple", 9.9)
		mock.ExpectQuery("SELECT id, name, price FROM products WHERE id = ?").
			WithArgs(1).
			WillReturnRows(rows)

		got, err := repo.FindByID(context.Background(), 1)
		require.NoError(t, err)
		assert.Equal(t, "Apple", got.Name)
		assert.Equal(t, 9.9, got.Price)
		require.NoError(t, mock.ExpectationsWereMet())
	})

	t.Run("not found", func(t *testing.T) {
		mock.ExpectQuery("SELECT id, name, price FROM products WHERE id = ?").
			WithArgs(999).
			WillReturnError(sql.ErrNoRows)

		_, err := repo.FindByID(context.Background(), 999)
		require.Error(t, err)
		assert.Contains(t, err.Error(), "not found")
		require.NoError(t, mock.ExpectationsWereMet())
	})

	t.Run("db error", func(t *testing.T) {
		mock.ExpectQuery("SELECT id, name, price FROM products WHERE id = ?").
			WithArgs(2).
			WillReturnError(errors.New("connection lost"))

		_, err := repo.FindByID(context.Background(), 2)
		require.Error(t, err)
		require.NoError(t, mock.ExpectationsWereMet())
	})
}

func TestMySQLRepo_Create(t *testing.T) {
	db, mock, err := sqlmock.New()
	require.NoError(t, err)
	defer db.Close()

	repo := product.NewMySQLRepo(db)

	mock.ExpectExec("INSERT INTO products").
		WithArgs("Apple", 9.9).
		WillReturnResult(sqlmock.NewResult(42, 1))

	got, err := repo.Create(context.Background(), product.Product{Name: "Apple", Price: 9.9})
	require.NoError(t, err)
	assert.Equal(t, 42, got.ID)
	require.NoError(t, mock.ExpectationsWereMet())
}

func TestMySQLRepo_List(t *testing.T) {
	db, mock, err := sqlmock.New()
	require.NoError(t, err)
	defer db.Close()

	repo := product.NewMySQLRepo(db)

	rows := sqlmock.NewRows([]string{"id", "name", "price"}).
		AddRow(1, "Apple", 9.9).
		AddRow(2, "Banana", 5.5)
	mock.ExpectQuery("SELECT id, name, price FROM products").
		WillReturnRows(rows)

	list, err := repo.List(context.Background())
	require.NoError(t, err)
	assert.Len(t, list, 2)
	assert.Equal(t, "Apple", list[0].Name)
	assert.Equal(t, "Banana", list[1].Name)
	require.NoError(t, mock.ExpectationsWereMet())
}

Service 层测试(基于 testify/mock)

service.go

go
package product

import (
	"context"
	"errors"
)

var (
	ErrInvalidProduct = errors.New("invalid product")
	ErrNotFound       = errors.New("not found")
)

type Service struct {
	repo Repository
}

func NewService(repo Repository) *Service {
	return &Service{repo: repo}
}

func (s *Service) Get(ctx context.Context, id int) (Product, error) {
	p, err := s.repo.FindByID(ctx, id)
	if err != nil {
		return Product{}, ErrNotFound
	}
	return p, nil
}

func (s *Service) Create(ctx context.Context, name string, price float64) (Product, error) {
	if name == "" || price <= 0 {
		return Product{}, ErrInvalidProduct
	}
	return s.repo.Create(ctx, Product{Name: name, Price: price})
}

func (s *Service) List(ctx context.Context) ([]Product, error) {
	return s.repo.List(ctx)
}

mock_repository.go(testify/mock 风格,通常用 mockery 生成):

go
package product

import (
	"context"

	"github.com/stretchr/testify/mock"
)

type MockRepository struct {
	mock.Mock
}

func (m *MockRepository) FindByID(ctx context.Context, id int) (Product, error) {
	args := m.Called(ctx, id)
	return args.Get(0).(Product), args.Error(1)
}

func (m *MockRepository) Create(ctx context.Context, p Product) (Product, error) {
	args := m.Called(ctx, p)
	return args.Get(0).(Product), args.Error(1)
}

func (m *MockRepository) List(ctx context.Context) ([]Product, error) {
	args := m.Called(ctx)
	if args.Get(0) == nil {
		return nil, args.Error(1)
	}
	return args.Get(0).([]Product), args.Error(1)
}

service_test.go

go
package product_test

import (
	"context"
	"errors"
	"testing"

	"github.com/stretchr/testify/assert"
	"github.com/stretchr/testify/mock"
	"github.com/stretchr/testify/require"

	"example.com/product"
)

func TestService_Get(t *testing.T) {
	t.Run("success", func(t *testing.T) {
		mockRepo := new(product.MockRepository)
		mockRepo.On("FindByID", mock.Anything, 1).
			Return(product.Product{ID: 1, Name: "Apple", Price: 9.9}, nil)

		svc := product.NewService(mockRepo)
		got, err := svc.Get(context.Background(), 1)

		require.NoError(t, err)
		assert.Equal(t, "Apple", got.Name)
		mockRepo.AssertExpectations(t)
	})

	t.Run("not found", func(t *testing.T) {
		mockRepo := new(product.MockRepository)
		mockRepo.On("FindByID", mock.Anything, 999).
			Return(product.Product{}, errors.New("not found"))

		svc := product.NewService(mockRepo)
		_, err := svc.Get(context.Background(), 999)

		require.Error(t, err)
		assert.ErrorIs(t, err, product.ErrNotFound)
		mockRepo.AssertExpectations(t)
	})
}

func TestService_Create(t *testing.T) {
	t.Run("valid", func(t *testing.T) {
		mockRepo := new(product.MockRepository)
		mockRepo.On("Create", mock.Anything, product.Product{Name: "Apple", Price: 9.9}).
			Return(product.Product{ID: 1, Name: "Apple", Price: 9.9}, nil)

		svc := product.NewService(mockRepo)
		got, err := svc.Create(context.Background(), "Apple", 9.9)

		require.NoError(t, err)
		assert.Equal(t, 1, got.ID)
		mockRepo.AssertExpectations(t)
	})

	t.Run("invalid name", func(t *testing.T) {
		mockRepo := new(product.MockRepository)
		svc := product.NewService(mockRepo)

		_, err := svc.Create(context.Background(), "", 9.9)
		require.Error(t, err)
		assert.ErrorIs(t, err, product.ErrInvalidProduct)
		// 不应该调用 repo.Create
		mockRepo.AssertNotCalled(t, "Create")
	})

	t.Run("invalid price", func(t *testing.T) {
		mockRepo := new(product.MockRepository)
		svc := product.NewService(mockRepo)

		_, err := svc.Create(context.Background(), "Apple", -1)
		require.Error(t, err)
		assert.ErrorIs(t, err, product.ErrInvalidProduct)
		mockRepo.AssertNotCalled(t, "Create")
	})
}

func TestService_List(t *testing.T) {
	t.Run("non empty", func(t *testing.T) {
		mockRepo := new(product.MockRepository)
		mockRepo.On("List", mock.Anything).
			Return([]product.Product{
				{ID: 1, Name: "Apple", Price: 9.9},
				{ID: 2, Name: "Banana", Price: 5.5},
			}, nil)

		svc := product.NewService(mockRepo)
		list, err := svc.List(context.Background())

		require.NoError(t, err)
		assert.Len(t, list, 2)
		mockRepo.AssertExpectations(t)
	})

	t.Run("empty", func(t *testing.T) {
		mockRepo := new(product.MockRepository)
		mockRepo.On("List", mock.Anything).
			Return([]product.Product{}, nil)

		svc := product.NewService(mockRepo)
		list, err := svc.List(context.Background())

		require.NoError(t, err)
		assert.Empty(t, list)
		mockRepo.AssertExpectations(t)
	})

	t.Run("nil slice", func(t *testing.T) {
		mockRepo := new(product.MockRepository)
		mockRepo.On("List", mock.Anything).
			Return(nil, nil)

		svc := product.NewService(mockRepo)
		list, err := svc.List(context.Background())

		require.NoError(t, err)
		assert.Empty(t, list)
		mockRepo.AssertExpectations(t)
	})
}

注意 List 在 Mock 里需要兼容 nil 返回值——上面 mock_repository.goList 实现已经处理了这种情况。

十、Mock 的常见陷阱

1. Mock 过度

不是所有依赖都需要 Mock。比如纯函数 Add(a, b int) 没有外部依赖,直接测就行。强行给纯函数包一层接口再 Mock,只会徒增复杂度。

2. Mock 反射真实行为

如果 Mock 的行为与真实依赖不一致,测试通过毫无意义。比如真实数据库在插入重复 key 时返回 ErrDuplicate,但 Mock 永远返回 nil,这样的测试会掩盖真实 bug。

3. 过度断言调用次数

Times(1)AssertNumberOfCalls(t, "X", 1) 这类严格断言会让测试变得脆弱。重构一个内部调用,多个测试就崩。只在「调用次数是契约一部分」时才断言次数,否则用 AnyTimes()

4. Mock 业务逻辑

不要在 Mock 里写复杂逻辑——比如「如果是 id=1 返回 Alice,如果是 id=2 返回 Bob」。如果 Mock 需要这么多分支,应该改用真实实现或测试替身。

5. 测试与 Mock 实现耦合过深

测试断言「Update 被以恰好这个参数调用」,往往把内部实现细节暴露给了测试。理想的测试应该断言「最终状态」或「外部可见行为」,而不是「内部调用了什么」。

十一、小结

本篇系统讲解了 Go 中 Mock 与接口测试的方方面面,核心要点:

  1. 依赖注入是可测试性的基础:依赖接口而非具体实现,测试时才能替换。
  2. 手动 Mock:最简单直接,适合接口简单、复用性强的场景。
  3. gomock:Google 出品,代码生成 + 严格校验,适合接口多、需要校验调用顺序的项目。
  4. mockery:配置驱动的生成工具,支持泛型,与 testify/mock 兼容。
  5. testify/mock:testify 自带的 Mock 框架,无需代码生成(也可生成),API 简洁,社区普及度高。
  6. 数据库 Mock:用 Repository 接口封装、用 sqlmock 模拟 sql.DB、用 SQLite 内存库做集成测试,三种方案各有所长。
  7. HTTP Mock:把 HTTP 调用抽象成接口,或用 httptest.NewServer 模拟服务器。
  8. Mock 陷阱:Mock 过度、反射失真、过度断言、Mock 业务逻辑、测试耦合过深。

下一篇我们将深入 HTTP 测试,使用标准库 httptest 包测试 HTTP Handler、中间件、客户端代码,并以 Gin 应用为例演示完整的 HTTP 测试套件。


记住一条原则:Mock 是手段,不是目的。Mock 的目的是让单元测试专注、快速、稳定。如果一份 Mock 让你维护测试的时间比维护生产代码还多,那很可能哪里出问题了。重新审视依赖边界,让接口更小、Mock 更简单,往往能从根本上解决问题。