Appearance
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 变得极其自然:
- 在被测代码中定义依赖接口。
- 在生产代码中实现该接口(如 MySQLRepo)。
- 在测试代码中写一个 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 生成工具,由两部分组成:
mockgen:命令行工具,根据接口定义自动生成 Mock 代码。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.go 的 List 实现已经处理了这种情况。
十、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 与接口测试的方方面面,核心要点:
- 依赖注入是可测试性的基础:依赖接口而非具体实现,测试时才能替换。
- 手动 Mock:最简单直接,适合接口简单、复用性强的场景。
- gomock:Google 出品,代码生成 + 严格校验,适合接口多、需要校验调用顺序的项目。
- mockery:配置驱动的生成工具,支持泛型,与 testify/mock 兼容。
- testify/mock:testify 自带的 Mock 框架,无需代码生成(也可生成),API 简洁,社区普及度高。
- 数据库 Mock:用 Repository 接口封装、用 sqlmock 模拟 sql.DB、用 SQLite 内存库做集成测试,三种方案各有所长。
- HTTP Mock:把 HTTP 调用抽象成接口,或用 httptest.NewServer 模拟服务器。
- Mock 陷阱:Mock 过度、反射失真、过度断言、Mock 业务逻辑、测试耦合过深。
下一篇我们将深入 HTTP 测试,使用标准库 httptest 包测试 HTTP Handler、中间件、客户端代码,并以 Gin 应用为例演示完整的 HTTP 测试套件。
记住一条原则:Mock 是手段,不是目的。Mock 的目的是让单元测试专注、快速、稳定。如果一份 Mock 让你维护测试的时间比维护生产代码还多,那很可能哪里出问题了。重新审视依赖边界,让接口更小、Mock 更简单,往往能从根本上解决问题。