Python 单元测试:给代码请"质检员"
引言:从"凭感觉"到"有标准"
想象你开了一家蛋糕店:
- 凭感觉:每次做完蛋糕尝一口,觉得好吃就卖——今天状态好,明天状态差,品质不稳定;
- 有标准(单元测试):制定检查清单——甜度 7 分、松软度 8 分、奶油厚度 2cm——每块蛋糕按清单逐项检查,合格才上架。
单元测试就是代码的"质检清单"——每次改代码,跑一遍清单,确保没把原来的功能改坏。
一、什么是单元测试?
1.1 定义
对一个模块、一个函数、一个类进行正确性检验的测试工作。
1.2 以 abs() 为例
测试用例设计:
| 输入 | 期待输出 | 测试点 |
|---|---|---|
1, 1.2, 0.99 | 与输入相同 | 正数 |
-1, -1.2, -0.99 | 与输入相反 | 负数 |
0 | 0 | 零 |
None, [], {} | TypeError | 非数值类型 |
核心思想:覆盖正常情况 + 边界情况 + 异常情况。
1.3 单元测试的价值
修改代码前:跑测试 → 全部通过
↓
修改代码
↓
再跑测试 → 通过?说明没改坏!不通过?说明改出问题了!生活化理解:单元测试是"存档点"——改代码前存档,改崩了读档重来,改对了继续前进。
二、实战:给 Dict 类写单元测试
2.1 被测代码:mydict.py
class Dict(dict):
"""支持属性访问的字典"""
def __init__(self, **kw):
super().__init__(**kw)
def __getattr__(self, key):
try:
return self[key]
except KeyError:
raise AttributeError("'Dict' object has no attribute '%s'" % key)
def __setattr__(self, key, value):
self[key] = value功能:d = Dict(a=1),既能 d['a'],也能 d.a。
2.2 测试代码:mydict_test.py
import unittest
from mydict import Dict
class TestDict(unittest.TestCase):
def test_init(self):
"""测试初始化"""
d = Dict(a=1, b='test')
self.assertEqual(d.a, 1)
self.assertEqual(d.b, 'test')
self.assertTrue(isinstance(d, dict))
def test_key(self):
"""测试 key 访问"""
d = Dict()
d['key'] = 'value'
self.assertEqual(d.key, 'value')
def test_attr(self):
"""测试属性访问"""
d = Dict()
d.key = 'value'
self.assertTrue('key' in d)
self.assertEqual(d['key'], 'value')
def test_keyerror(self):
"""测试 key 不存在时抛 KeyError"""
d = Dict()
with self.assertRaises(KeyError):
value = d['empty']
def test_attrerror(self):
"""测试属性不存在时抛 AttributeError"""
d = Dict()
with self.assertRaises(AttributeError):
value = d.empty2.3 关键规则
| 规则 | 说明 |
|---|---|
继承 unittest.TestCase | 测试类的"上岗证" |
方法名以 test_ 开头 | 只有 test_xxx 会被执行 |
用 self.assertXxx() | 不用 print,用断言自动判断 |
2.4 常用断言方法
self.assertEqual(a, b) # a == b
self.assertTrue(x) # bool(x) is True
self.assertFalse(x) # bool(x) is False
self.assertIsNone(x) # x is None
self.assertIn(a, b) # a in b
# 期待抛出异常
with self.assertRaises(KeyError):
d['empty']三、运行单元测试
3.1 方式一:脚本直接运行
# mydict_test.py 末尾加
if __name__ == '__main__':
unittest.main()python mydict_test.py3.2 方式二:命令行运行(推荐)
python -m unittest mydict_test输出:
.....
----------------------------------------------------------------------
Ran 5 tests in 0.000s
OK推荐原因:可以批量运行多个测试文件,很多工具能自动执行。
3.3 运行单个测试
python -m unittest mydict_test.TestDict.test_attr3.4 按关键词筛选
python -m unittest mydict_test -k attr -v-k attr:匹配含 "attr" 的测试方法;-v:显示详细信息。
四、setUp 与 tearDown:测试前后的"准备和收拾"
4.1 场景
每个测试都需要连接数据库,测试完要断开——重复代码写在哪?
4.2 解决方案
class TestDict(unittest.TestCase):
def setUp(self):
"""每个测试方法执行前调用"""
print('setUp...')
self.db = connect_database() # 连接数据库
def tearDown(self):
"""每个测试方法执行后调用"""
print('tearDown...')
self.db.close() # 关闭连接
def test_query(self):
result = self.db.query('SELECT * FROM users')
self.assertIsNotNone(result)执行顺序:
setUp → test_query → tearDown
setUp → test_insert → tearDown
setUp → test_delete → tearDown生活化理解:setUp 是"做饭前洗菜、备料",tearDown 是"吃完饭洗碗、擦桌子"——每顿饭(每个测试)都要做,但不用写进每道菜(每个测试方法)的做法里。
五、知识链条:从写代码到改代码
编写功能代码(Dict 类)
↓
设计测试用例(正常 + 边界 + 异常)
↓
编写测试代码(TestCase + test_xxx + assertXxx)
↓
运行测试(python -m unittest)
↓
全部通过? → 功能正确,可以提交
↓
需要改代码? → 改完再跑测试 → 通过?没改坏! → 继续开发六、常见误区与避坑指南
6.1 误区一:测试代码太复杂
def test_complex(self):
# ❌ 测试里又写了一套业务逻辑
for i in range(100):
for j in range(100):
if i % 2 == 0 and j % 3 == 0:
self.assertEqual(complex_calc(i, j), expected[i][j])问题:测试代码本身可能有 bug,谁又来测测试?
原则:测试代码要非常简单——简单到一眼就能看出对不对。
6.2 误区二:只测正常情况
def test_add(self):
self.assertEqual(add(1, 2), 3) # 只测了正数修正:覆盖三类情况:
def test_add_positive(self):
self.assertEqual(add(1, 2), 3)
def test_add_negative(self):
self.assertEqual(add(-1, -2), -3)
def test_add_zero(self):
self.assertEqual(add(0, 0), 0)6.3 误区三:测试依赖外部环境
def test_download(self):
result = download('https://example.com/file.zip') # ❌ 依赖网络
self.assertTrue(len(result) > 0)问题:网络不通就失败,但代码本身没问题。
修正:用 mock 模拟外部依赖,或标记为集成测试(不属于单元测试)。
6.4 误区四:认为测试通过 = 没有 bug
单元测试通过 ≠ 程序没有 bug,但不通过 = 肯定有 bug。
测试只能证明"在测试用例覆盖的范围内,行为正确"——范围外的 bug 测不出来。
七、实际应用案例
案例 1:用户注册验证器
# validator.py
class ValidationError(ValueError):
pass
def validate_email(email):
if '@' not in email:
raise ValidationError('邮箱必须包含 @')
if '.' not in email.split('@')[-1]:
raise ValidationError('邮箱域名必须包含 .')
return True
def validate_age(age):
if not isinstance(age, int):
raise ValidationError('年龄必须是整数')
if age < 0 or age > 150:
raise ValidationError('年龄必须在 0-150 之间')
return True# test_validator.py
import unittest
from validator import validate_email, validate_age, ValidationError
class TestValidator(unittest.TestCase):
# 正常情况
def test_valid_email(self):
self.assertTrue(validate_email('user@example.com'))
def test_valid_age(self):
self.assertTrue(validate_age(25))
# 边界情况
def test_email_min_length(self):
self.assertTrue(validate_email('a@b.c'))
def test_age_boundary(self):
self.assertTrue(validate_age(0))
self.assertTrue(validate_age(150))
# 异常情况
def test_email_no_at(self):
with self.assertRaises(ValidationError):
validate_email('userexample.com')
def test_email_no_dot(self):
with self.assertRaises(ValidationError):
validate_email('user@example')
def test_age_negative(self):
with self.assertRaises(ValidationError):
validate_age(-1)
def test_age_too_large(self):
with self.assertRaises(ValidationError):
validate_age(151)
def test_age_not_int(self):
with self.assertRaises(ValidationError):
validate_age('25')
if __name__ == '__main__':
unittest.main()生活化理解:验证器是"机场安检"——邮箱、年龄是"旅客",测试用例是"各种证件组合":正常护照、过期护照、假护照、没带护照——每种情况都要确保安检系统正确处理。
案例 2:购物车计算
# cart.py
class Cart(object):
def __init__(self):
self.items = []
def add(self, name, price, count=1):
if price < 0:
raise ValueError('价格不能为负')
if count <= 0:
raise ValueError('数量必须为正数')
self.items.append({'name': name, 'price': price, 'count': count})
def total(self):
return sum(item['price'] * item['count'] for item in self.items)
def discount(self, rate):
"""打折,rate 为折扣率,如 0.8 表示 8 折"""
if not 0 < rate <= 1:
raise ValueError('折扣率必须在 0-1 之间')
return self.total() * rate# test_cart.py
import unittest
from cart import Cart
class TestCart(unittest.TestCase):
def setUp(self):
"""每个测试前创建新购物车"""
self.cart = Cart()
def test_empty_cart(self):
"""空购物车总价为 0"""
self.assertEqual(self.cart.total(), 0)
def test_add_single_item(self):
self.cart.add('苹果', 5, 2)
self.assertEqual(self.cart.total(), 10)
def test_add_multiple_items(self):
self.cart.add('苹果', 5, 2)
self.cart.add('香蕉', 3, 3)
self.assertEqual(self.cart.total(), 19)
def test_discount(self):
self.cart.add('苹果', 10, 1)
self.assertEqual(self.cart.discount(0.8), 8.0)
def test_invalid_price(self):
with self.assertRaises(ValueError):
self.cart.add('苹果', -5, 1)
def test_invalid_count(self):
with self.assertRaises(ValueError):
self.cart.add('苹果', 5, 0)
def test_invalid_discount(self):
self.cart.add('苹果', 10, 1)
with self.assertRaises(ValueError):
self.cart.discount(1.5) # 折扣率不能大于 1
if __name__ == '__main__':
unittest.main()setUp 的价值:每个测试都用全新的 Cart,避免测试间互相影响。
八、实战练习
练习:修复 Student.get_grade() 的 bug
下面代码有 bug,测试不通过,请修复 Student 类:
import unittest
class Student(object):
def __init__(self, name, score):
self.name = name
self.score = score
def get_grade(self):
if self.score >= 60:
return 'B'
if self.score >= 80:
return 'A'
return 'C'
class TestStudent(unittest.TestCase):
def test_80_to_100(self):
s1 = Student('Bart', 80)
s2 = Student('Lisa', 100)
self.assertEqual(s1.get_grade(), 'A')
self.assertEqual(s2.get_grade(), 'A')
def test_60_to_80(self):
s1 = Student('Bart', 60)
s2 = Student('Lisa', 79)
self.assertEqual(s1.get_grade(), 'B')
self.assertEqual(s2.get_grade(), 'B')
def test_0_to_60(self):
s1 = Student('Bart', 0)
s2 = Student('Lisa', 59)
self.assertEqual(s1.get_grade(), 'C')
self.assertEqual(s2.get_grade(), 'C')
def test_invalid(self):
s1 = Student('Bart', -1)
s2 = Student('Lisa', 101)
with self.assertRaises(ValueError):
s1.get_grade()
with self.assertRaises(ValueError):
s2.get_grade()
if __name__ == '__main__':
unittest.main()参考答案
Bug 分析:if self.score >= 60 先执行,80 分也会进入这个分支返回 'B',永远到不了 'A'。
修复:调整判断顺序,高分在前:
class Student(object):
def __init__(self, name, score):
self.name = name
self.score = score
def get_grade(self):
if self.score < 0 or self.score > 100:
raise ValueError('分数必须在 0-100 之间')
if self.score >= 80:
return 'A'
if self.score >= 60:
return 'B'
return 'C'测试通过:
....
----------------------------------------------------------------------
Ran 4 tests in 0.000s
OK九、小结
- 单元测试:对模块/函数/类的正确性检验,覆盖正常 + 边界 + 异常;
- 核心结构:继承
unittest.TestCase,方法名test_xxx,用self.assertXxx()断言; - 运行方式:
python -m unittest test_file,支持单测、筛选、详细输出; setUp/tearDown:每个测试前后的准备和清理,避免重复代码;- 测试原则:测试代码要简单,简单到一眼看出对错;
- 价值认知:测试通过 ≠ 没 bug,但不通过 = 肯定有 bug;
- 终极价值:改代码的信心保证——有测试护航,重构不心慌。
单元测试是代码的"保险单"——平时觉得多余,出事时才知道值。