Skip to content

Python 单元测试:给代码请"质检员"

引言:从"凭感觉"到"有标准"

想象你开了一家蛋糕店:

  • 凭感觉:每次做完蛋糕尝一口,觉得好吃就卖——今天状态好,明天状态差,品质不稳定;
  • 有标准(单元测试):制定检查清单——甜度 7 分、松软度 8 分、奶油厚度 2cm——每块蛋糕按清单逐项检查,合格才上架。

单元测试就是代码的"质检清单"——每次改代码,跑一遍清单,确保没把原来的功能改坏


一、什么是单元测试?

1.1 定义

一个模块、一个函数、一个类进行正确性检验的测试工作。

1.2 以 abs() 为例

测试用例设计:

输入期待输出测试点
1, 1.2, 0.99与输入相同正数
-1, -1.2, -0.99与输入相反负数
00
None, [], {}TypeError非数值类型

核心思想:覆盖正常情况 + 边界情况 + 异常情况

1.3 单元测试的价值

修改代码前:跑测试 → 全部通过

修改代码

再跑测试 → 通过?说明没改坏!不通过?说明改出问题了!

生活化理解:单元测试是"存档点"——改代码前存档,改崩了读档重来,改对了继续前进。


二、实战:给 Dict 类写单元测试

2.1 被测代码:mydict.py

python
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

python
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.empty

2.3 关键规则

规则说明
继承 unittest.TestCase测试类的"上岗证"
方法名以 test_ 开头只有 test_xxx 会被执行
self.assertXxx()不用 print,用断言自动判断

2.4 常用断言方法

python
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 方式一:脚本直接运行

python
# mydict_test.py 末尾加
if __name__ == '__main__':
    unittest.main()
bash
python mydict_test.py

3.2 方式二:命令行运行(推荐)

bash
python -m unittest mydict_test

输出

.....
----------------------------------------------------------------------
Ran 5 tests in 0.000s

OK

推荐原因:可以批量运行多个测试文件,很多工具能自动执行。

3.3 运行单个测试

bash
python -m unittest mydict_test.TestDict.test_attr

3.4 按关键词筛选

bash
python -m unittest mydict_test -k attr -v

-k attr:匹配含 "attr" 的测试方法;-v:显示详细信息。


四、setUptearDown:测试前后的"准备和收拾"

4.1 场景

每个测试都需要连接数据库,测试完要断开——重复代码写在哪

4.2 解决方案

python
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 误区一:测试代码太复杂

python
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 误区二:只测正常情况

python
def test_add(self):
    self.assertEqual(add(1, 2), 3)   # 只测了正数

修正:覆盖三类情况:

python
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 误区三:测试依赖外部环境

python
def test_download(self):
    result = download('https://example.com/file.zip')   # ❌ 依赖网络
    self.assertTrue(len(result) > 0)

问题:网络不通就失败,但代码本身没问题。

修正:用 mock 模拟外部依赖,或标记为集成测试(不属于单元测试)。

6.4 误区四:认为测试通过 = 没有 bug

单元测试通过 ≠ 程序没有 bug,但不通过 = 肯定有 bug。

测试只能证明"在测试用例覆盖的范围内,行为正确"——范围外的 bug 测不出来。


七、实际应用案例

案例 1:用户注册验证器

python
# 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
python
# 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:购物车计算

python
# 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
python
# 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 类:

python
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'

修复:调整判断顺序,高分在前:

python
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

九、小结

  1. 单元测试:对模块/函数/类的正确性检验,覆盖正常 + 边界 + 异常
  2. 核心结构:继承 unittest.TestCase,方法名 test_xxx,用 self.assertXxx() 断言;
  3. 运行方式python -m unittest test_file,支持单测、筛选、详细输出;
  4. setUp / tearDown:每个测试前后的准备和清理,避免重复代码;
  5. 测试原则测试代码要简单,简单到一眼看出对错;
  6. 价值认知:测试通过 ≠ 没 bug,但不通过 = 肯定有 bug
  7. 终极价值改代码的信心保证——有测试护航,重构不心慌。

单元测试是代码的"保险单"——平时觉得多余,出事时才知道值。