单元测试
测试是程序员的"安全网"。Python 标准库自带 unittest,而社区主流是 pytest(本沙箱未安装,需要 pip install pytest,但 API 风格值得了解)。
1. 第一个测试:unittest
每个测试都是一个方法名以 test_ 开头的方法。
import unittest
def add(a, b):
return a + b
def divide(a, b):
if b == 0:
raise ValueError("除数不能为 0")
return a / b
class TestMath(unittest.TestCase):
def test_add(self):
self.assertEqual(add(3, 4), 7)
self.assertEqual(add(-1, 1), 0)
def test_divide_ok(self):
self.assertEqual(divide(10, 2), 5)
def test_divide_by_zero(self):
with self.assertRaises(ValueError):
divide(1, 0)
# 把 unittest 当主程序跑
unittest.main(argv=[''], exit=False, verbosity=2)ℹ️常用的断言方法
assertEqual(a, b)/assertNotEqualassertTrue(x)/assertFalse(x)assertIsNone(x)/assertIsNotNoneassertIn(a, b)/assertNotInassertRaises(Exc)配合withassertAlmostEqual(a, b, places=4)比较浮点
2. 生命周期:setUp / tearDown
每个测试方法前/后自动调用:
import unittest
import tempfile, os
class FileTest(unittest.TestCase):
def setUp(self):
# 每个测试前建一个临时文件
self.path = tempfile.mktemp(suffix=".txt")
with open(self.path, "w") as f:
f.write("hello\nworld\n")
def tearDown(self):
# 每个测试后清理
if os.path.exists(self.path):
os.remove(self.path)
def test_read_line_count(self):
with open(self.path) as f:
lines = f.readlines()
self.assertEqual(len(lines), 2)
def test_contains_hello(self):
with open(self.path) as f:
content = f.read()
self.assertIn("hello", content)
unittest.main(argv=[''], exit=False, verbosity=2)3. 单纯 assert 风格(pytest 风格)
pytest 之所以流行,是因为它允许直接用 Python 的 assert 语句,失败时打印出具体的中间值。这非常适合用普通函数写(不必继承类)。本节展示风格示例——运行请在自己环境 pip install pytest 后 pytest test_xxx.py。
# 保存为 test_sample.py,然后用 pytest 运行
def add(a, b):
return a + b
def test_add_positive():
assert add(3, 4) == 7
def test_add_negative():
assert add(-1, 1) == 0
def test_add_strings():
# pytest 会显示 "hi" + "" != "hi!"
assert add("hi", "") == "hi!"💡为什么 pytest 流行
- 无需继承
TestCase assert失败时打印完整的上下文表达式(不只是 True/False)fixture机制强大,参数化一行搞定- 大量插件(pytest-cov、pytest-mock、pytest-asyncio)
4. 用 subTest 做参数化(stdlib 方案)
不依赖 pytest,也能参数化:
import unittest
def slugify(s):
return s.lower().strip().replace(" ", "-")
class TestSlugify(unittest.TestCase):
def test_cases(self):
cases = [
("Hello World", "hello-world"),
(" Python ", "python"),
("ALREADY-low", "already-low"),
("a b c d", "a-b-c-d"),
]
for input_, expected in cases:
with self.subTest(input=input_):
self.assertEqual(slugify(input_), expected)
unittest.main(argv=[''], exit=False, verbosity=2)ℹ️pytest 的 parametrize 等价
import pytest
@pytest.mark.parametrize("input,expected", [
("Hello World", "hello-world"),
(" Python ", "python"),
])
def test_slugify(input, expected):
assert slugify(input) == expected失败时会一个 case 一个 case 报告。
5. Mock —— 替换掉不可控的依赖
测试时不想真发邮件、查数据库、调网络?用 unittest.mock 替换掉。
import unittest
from unittest.mock import patch, MagicMock
# 被测代码:从一个 dict 拿用户
def get_user_email(uid, db):
user = db.fetch_user(uid)
if not user:
return None
return user["email"]
class TestGetUserEmail(unittest.TestCase):
@patch("__main__.db")
def test_existing_user(self, mock_db):
# 让 db.fetch_user 返回我们假装的数据
mock_db.fetch_user.return_value = {"name": "Alice", "email": "a@x.com"}
self.assertEqual(get_user_email(1, mock_db), "a@x.com")
mock_db.fetch_user.assert_called_once_with(1)
@patch("__main__.db")
def test_missing_user(self, mock_db):
mock_db.fetch_user.return_value = None
self.assertIsNone(get_user_email(999, mock_db))
# 假装有个 db 对象
class FakeDB: pass
db = FakeDB()
unittest.main(argv=[''], exit=False, verbosity=2)💡Mock 三件套
@patch("module.func"):测试期间把func替换成 MagicMockmock.return_value:调用时的返回值mock.side_effect = [1, 2, ValueError]:每次调用依次返回,抛异常时停
6. 进阶:模拟 time.sleep
import unittest
from unittest.mock import patch
import time
# 被测代码
def wait_then_print():
time.sleep(10) # 真跑要 10 秒
print("done!")
class TestWait(unittest.TestCase):
@patch("time.sleep")
def test_does_not_sleep(self, mock_sleep):
wait_then_print()
mock_sleep.assert_called_once_with(10)
unittest.main(argv=[''], exit=False, verbosity=2)🎯 练习
为下面这个 clamp(x, lo, hi) 函数(把 x 限制在 [lo, hi] 区间内)写至少 3 个测试:① 在区间内;② 低于下界;③ 高于上界;④ lo == hi。
import unittest
def clamp(x, lo, hi):
if x < lo: return lo
if x > hi: return hi
return x
class TestClamp(unittest.TestCase):
def test_inside(self):
pass
def test_below(self):
pass
def test_above(self):
pass
def test_equal(self):
pass
unittest.main(argv=[''], exit=False, verbosity=2)🎯提示
inside:clamp(5, 0, 10) == 5below:clamp(-3, 0, 10) == 0above:clamp(99, 0, 10) == 10equal:clamp(5, 5, 5) == 5
小结
- ✅
unittest.TestCase+ 方法名test_xxx是 stdlib 写法 - ✅
assertEqual/assertTrue/assertRaises是常用断言 - ✅
setUp/tearDown在每个测试前后自动跑 - ✅ pytest 风格:
assert+ 上下文输出 + 插件生态(需pip install pytest) - ✅
unittest.mock.patch替换依赖;return_value/side_effect控制行为 - ✅
subTest在 stdlib 里做参数化;pytest 用@parametrize
下一章 常用第三方库:requests、pandas 等怎么用、为什么用、装不上时怎么办。