1. Python单元测试(unittest)实战指南
单元测试是软件开发中不可或缺的一环,它能确保代码的每个独立部分按预期工作。Python内置的unittest模块为我们提供了强大的测试框架,但很多开发者对其使用仍停留在基础层面。本文将带你深入unittest的实战应用,从基础到高级技巧,让你写出更健壮、可维护的测试代码。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. unittest基础入门
2.1 unittest核心组件解析
unittest框架包含几个关键组件:
- TestCase:测试用例基类,所有测试类都应继承此类
- TestSuite:测试套件,用于组织多个测试用例
- TestLoader:用于从类和模块中发现测试
- TextTestRunner:运行测试并输出结果的运行器
- TestResult:收集测试结果的对象
一个最基本的测试用例示例如下:
python复制import unittest
class TestStringMethods(unittest.TestCase):
def test_upper(self):
self.assertEqual('foo'.upper(), 'FOO')
def test_isupper(self):
self.assertTrue('FOO'.isupper())
self.assertFalse('Foo'.isupper())
if __name__ == '__main__':
unittest.main()
2.2 常用断言方法详解
unittest提供了丰富的断言方法,以下是常用的几种:
| 断言方法 | 检查条件 | 示例 |
|---|---|---|
| assertEqual(a, b) | a == b | self.assertEqual(1+1, 2) |
| assertNotEqual(a, b) | a != b | self.assertNotEqual(1, 2) |
| assertTrue(x) | bool(x) is True | self.assertTrue(1 < 2) |
| assertFalse(x) | bool(x) is False | self.assertFalse(1 > 2) |
| assertIs(a, b) | a is b | self.assertIs(None, None) |
| assertIsNot(a, b) | a is not b | self.assertIsNot([], []) |
| assertIsNone(x) | x is None | self.assertIsNone(None) |
| assertIsNotNone(x) | x is not None | self.assertIsNotNone(1) |
| assertIn(a, b) | a in b | self.assertIn('a', 'abc') |
| assertNotIn(a, b) | a not in b | self.assertNotIn('d', 'abc') |
| assertIsInstance(a, b) | isinstance(a, b) | self.assertIsInstance(1, int) |
| assertNotIsInstance(a, b) | not isinstance(a, b) | self.assertNotIsInstance(1, str) |
提示:优先使用最具体的断言方法,比如检查None时用assertIsNone而不是assertEqual(x, None),这样测试失败时的错误信息会更明确。
3. 高级测试技巧
3.1 测试固件(Fixture)的使用
测试固件是指在测试运行前后执行的代码,用于准备和清理测试环境。unittest提供了setUp()和tearDown()方法:
python复制class DatabaseTest(unittest.TestCase):
def setUp(self):
self.conn = create_db_connection()
self.cur = self.conn.cursor()
def tearDown(self):
self.cur.close()
self.conn.close()
def test_query(self):
self.cur.execute("SELECT 1")
result = self.cur.fetchone()
self.assertEqual(result, (1,))
此外,还有类级别的固件@classmethod setUpClass和tearDownClass,它们在整个测试类运行前后各执行一次。
3.2 参数化测试的实现
unittest本身不直接支持参数化测试,但可以通过以下方式实现:
- 使用subTest上下文管理器(Python 3.4+):
python复制class TestNumbers(unittest.TestCase):
def test_even(self):
for i in range(0, 6):
with self.subTest(i=i):
self.assertEqual(i % 2, 0)
- 使用第三方库如parameterized:
python复制from parameterized import parameterized
class TestMath(unittest.TestCase):
@parameterized.expand([
("negative", -1.5, -2.0),
("integer", 1, 1.0),
("large fraction", 1.6, 1),
])
def test_floor(self, name, input, expected):
self.assertEqual(math.floor(input), expected)
3.3 模拟对象(Mock)的使用
unittest.mock模块提供了Mock和MagicMock类,用于创建测试替身:
python复制from unittest.mock import Mock, patch
class TestPayment(unittest.TestCase):
def test_payment_success(self):
payment_gateway = Mock()
payment_gateway.process.return_value = True
result = process_payment(payment_gateway, 100)
self.assertTrue(result)
payment_gateway.process.assert_called_once_with(100)
@patch('module.ExternalAPI')
def test_api_call(self, mock_api):
mock_api.return_value.get_data.return_value = {'key': 'value'}
result = call_external_api()
self.assertEqual(result, {'key': 'value'})
注意:使用patch装饰器时,必须确保在正确的命名空间打补丁。常见的错误是在测试模块中导入被测试函数使用的对象,而不是在函数内部导入。
4. 测试组织与运行
4.1 测试发现与执行
unittest支持自动发现测试:
bash复制# 发现并运行当前目录下所有test_*.py文件中的测试
python -m unittest discover
# 指定测试目录
python -m unittest discover -s project/tests
# 运行特定测试模块
python -m unittest test_module1 test_module2
# 运行特定测试类
python -m unittest test_module.TestClass
# 运行特定测试方法
python -m unittest test_module.TestClass.test_method
4.2 测试套件的自定义
对于大型项目,可能需要自定义测试套件:
python复制def suite():
suite = unittest.TestSuite()
suite.addTest(TestStringMethods('test_upper'))
suite.addTest(TestStringMethods('test_isupper'))
# 或者添加整个测试类
suite.addTest(unittest.makeSuite(TestDatabase))
return suite
if __name__ == '__main__':
runner = unittest.TextTestRunner()
runner.run(suite())
4.3 测试覆盖率统计
使用coverage.py工具可以统计测试覆盖率:
bash复制# 安装
pip install coverage
# 运行测试并收集覆盖率数据
coverage run -m unittest discover
# 生成报告
coverage report -m
coverage html # 生成HTML报告
5. 实战经验与常见问题
5.1 测试金字塔原则
遵循测试金字塔原则:
- 大量小而快的单元测试(底层)
- 适量集成测试(中层)
- 少量端到端测试(顶层)
单元测试应该:
- 运行速度快(毫秒级)
- 相互独立
- 不依赖外部环境
- 测试单一功能点
5.2 测试命名最佳实践
好的测试名称应该:
- 描述被测试的功能
- 描述测试场景
- 描述预期结果
例如:
- test_addition_with_positive_numbers
- test_login_with_invalid_credentials_should_fail
- test_user_creation_with_missing_required_field
避免使用test1、test2这样的名称。
5.3 常见问题与解决方案
-
测试依赖外部服务:
- 使用Mock替代真实服务
- 考虑使用测试专用数据库
- 对于不可避免的外部依赖,使用@unittest.skipIf标记
-
测试随机失败:
- 确保测试不依赖执行顺序
- 避免共享状态
- 使用随机种子确保可重复性
-
测试运行缓慢:
- 减少I/O操作
- 使用内存数据库
- 并行运行测试(可使用pytest-xdist)
-
测试难以维护:
- 遵循DRY原则,提取公共代码
- 使用工厂函数创建测试数据
- 保持测试代码与生产代码同等质量
5.4 测试驱动开发(TDD)实践
TDD的基本流程:
- 编写一个失败的测试
- 编写最少代码使测试通过
- 重构代码,保持测试通过
示例:
python复制# 第一步:编写测试
class TestCalculator(unittest.TestCase):
def test_add(self):
calc = Calculator()
self.assertEqual(calc.add(2, 3), 5)
# 第二步:实现代码
class Calculator:
def add(self, a, b):
return a + b
# 第三步:重构(如果需要)
6. 与其他测试框架的比较
6.1 unittest vs pytest
| 特性 | unittest | pytest |
|---|---|---|
| 安装 | Python内置 | 需要安装 |
| 断言 | 使用assert*方法 | 直接使用assert语句 |
| 参数化 | 需要额外代码 | 内置支持 |
| 固件 | setUp/tearDown | 更灵活的fixture系统 |
| 插件 | 有限 | 丰富生态系统 |
| 测试发现 | 需要特定命名 | 更灵活的发现机制 |
6.2 何时选择unittest
- 项目要求使用标准库
- 需要与现有unittest测试套件集成
- 团队熟悉JUnit风格测试
- 需要与某些特定工具集成(如某些CI系统)
6.3 迁移到pytest的考虑
pytest可以运行unittest测试,因此可以逐步迁移:
- 先安装pytest
- 用pytest运行现有unittest测试
- 逐步将测试改写为pytest风格
- 利用pytest特有功能
7. 持续集成中的单元测试
7.1 与CI工具集成
主流CI工具都支持Python单元测试:
GitHub Actions示例:
yaml复制name: Python unittest
on: [push, pull_request]
jobs:
test:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- name: Set up Python
uses: actions/setup-python@v2
with:
python-version: '3.9'
- name: Install dependencies
run: |
python -m pip install --upgrade pip
pip install -r requirements.txt
- name: Run tests
run: |
python -m unittest discover
7.2 测试结果报告
生成JUnit格式报告便于CI工具解析:
bash复制python -m unittest discover -s tests -p "test_*.py" -v 2>&1 | tee test-results.txt
或者使用xmlrunner:
python复制import unittest
import xmlrunner
if __name__ == '__main__':
unittest.main(
testRunner=xmlrunner.XMLTestRunner(output='test-reports'),
failfast=False,
buffer=False,
catchbreak=False)
8. 性能测试与基准测试
虽然unittest主要用于功能测试,但也可以用于简单的性能检查:
python复制import time
import unittest
class TestPerformance(unittest.TestCase):
def test_fast_enough(self):
start = time.perf_counter()
result = expensive_operation()
duration = time.perf_counter() - start
self.assertLess(duration, 1.0) # 应在1秒内完成
self.assertEqual(result, expected_value)
对于更复杂的基准测试,可以考虑使用timeit模块或专门的基准测试框架如pytest-benchmark。
9. 测试私有方法
测试私有方法(以_开头)通常不被推荐,因为这违反了封装原则。更好的做法是:
- 通过公有方法测试私有方法的功能
- 如果必须测试,可以使用以下方式:
python复制import unittest
class TestPrivateMethods(unittest.TestCase):
def test_private_method(self):
obj = MyClass()
# 访问私有方法
result = obj._private_method()
self.assertEqual(result, expected_value)
或者使用inspect模块:
python复制import inspect
method = inspect.getattr_static(MyClass, '_private_method')
result = method(obj)
10. 测试异步代码
对于asyncio代码,可以使用unittest.IsolatedAsyncioTestCase(Python 3.8+):
python复制import asyncio
import unittest
class TestAsync(unittest.IsolatedAsyncioTestCase):
async def test_coroutine(self):
result = await async_function()
self.assertEqual(result, expected_value)
对于更早的Python版本,可以手动运行事件循环:
python复制class TestAsync(unittest.TestCase):
def test_coroutine(self):
loop = asyncio.get_event_loop()
result = loop.run_until_complete(async_function())
self.assertEqual(result, expected_value)
11. 数据库测试最佳实践
测试数据库相关代码时:
- 使用内存数据库(SQLite)加速测试
- 每个测试使用独立的事务并在测试后回滚
- 考虑使用工厂模式创建测试数据
示例:
python复制import sqlite3
import unittest
class TestDatabase(unittest.TestCase):
def setUp(self):
self.conn = sqlite3.connect(':memory:')
self.cur = self.conn.cursor()
self.cur.execute("CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT)")
def tearDown(self):
self.conn.close()
def test_insert(self):
self.cur.execute("INSERT INTO users (name) VALUES (?)", ('Alice',))
self.cur.execute("SELECT name FROM users WHERE id = 1")
result = self.cur.fetchone()
self.assertEqual(result, ('Alice',))
12. 测试异常处理
测试代码是否抛出预期异常:
python复制class TestExceptions(unittest.TestCase):
def test_division_by_zero(self):
with self.assertRaises(ZeroDivisionError):
1 / 0
def test_custom_exception(self):
with self.assertRaises(ValueError) as cm:
validate_input(-1)
self.assertEqual(str(cm.exception), "Input must be positive")
13. 测试随机性和不确定性代码
对于包含随机性的代码:
- 设置随机种子确保可重复性
- 测试统计属性而非具体值
- 多次运行测试验证稳定性
示例:
python复制import random
class TestRandom(unittest.TestCase):
def test_random_choice(self):
random.seed(42) # 固定随机种子
items = [1, 2, 3, 4]
choice = random.choice(items)
self.assertEqual(choice, 3) # 种子固定时结果可预测
def test_random_distribution(self):
counts = [0] * 10
for _ in range(10000):
num = random.randint(0, 9)
counts[num] += 1
# 每个数字出现次数应该在900-1100之间
for count in counts:
self.assertTrue(900 <= count <= 1100)
14. 测试日志输出
验证代码是否产生正确的日志:
python复制import logging
import unittest
from io import StringIO
class TestLogging(unittest.TestCase):
def test_log_output(self):
logger = logging.getLogger('test')
logger.setLevel(logging.INFO)
stream = StringIO()
handler = logging.StreamHandler(stream)
logger.addHandler(handler)
logger.info('Test message')
log_output = stream.getvalue()
self.assertIn('Test message', log_output)
15. 测试日期和时间相关代码
测试时间相关代码的挑战:
- 使用freezegun库冻结时间:
python复制from freezegun import freeze_time
import datetime
class TestTime(unittest.TestCase):
@freeze_time("2023-01-01")
def test_new_year(self):
now = datetime.datetime.now()
self.assertEqual(now.year, 2023)
self.assertEqual(now.month, 1)
self.assertEqual(now.day, 1)
- 注入时钟依赖:
python复制class TimeDependent:
def __init__(self, clock=datetime.datetime.now):
self.clock = clock
def is_weekend(self):
return self.clock().weekday() >= 5
class TestTimeDependent(unittest.TestCase):
def test_weekend(self):
fixed_time = lambda: datetime.datetime(2023, 1, 1) # 周日
td = TimeDependent(fixed_time)
self.assertTrue(td.is_weekend())
16. 测试文件操作
测试文件系统操作的最佳实践:
- 使用tempfile模块创建临时文件
- 使用unittest.mock.patch模拟文件操作
- 测试后清理临时文件
示例:
python复制import tempfile
import unittest
class TestFileOperations(unittest.TestCase):
def setUp(self):
self.temp_file = tempfile.NamedTemporaryFile(delete=False)
self.temp_path = self.temp_file.name
def tearDown(self):
import os
if os.path.exists(self.temp_path):
os.unlink(self.temp_path)
def test_write_file(self):
with open(self.temp_path, 'w') as f:
f.write('test content')
with open(self.temp_path) as f:
content = f.read()
self.assertEqual(content, 'test content')
17. 测试命令行接口
测试CLI应用程序:
- 使用subprocess运行命令并检查输出
- 使用unittest.mock.patch模拟sys.argv
- 使用click.testing或argparse.testing等专用测试工具
示例:
python复制import subprocess
import unittest
class TestCLI(unittest.TestCase):
def test_cli_output(self):
result = subprocess.run(
['python', 'script.py', '--version'],
capture_output=True,
text=True
)
self.assertEqual(result.returncode, 0)
self.assertIn('1.0.0', result.stdout)
18. 测试Web应用程序
测试Web应用的策略:
- 使用unittest.mock模拟HTTP请求
- 使用测试客户端(如Flask.test_client)
- 使用requests-mock库模拟requests
Flask应用测试示例:
python复制from flask import Flask
import unittest
app = Flask(__name__)
@app.route('/')
def index():
return 'Hello, World!'
class TestFlaskApp(unittest.TestCase):
def setUp(self):
self.app = app.test_client()
def test_index(self):
response = self.app.get('/')
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data.decode(), 'Hello, World!')
19. 测试Django应用
Django提供了扩展的测试工具:
python复制from django.test import TestCase
from myapp.models import Book
class BookTestCase(TestCase):
def setUp(self):
Book.objects.create(title="Test Book", author="Test Author")
def test_book_creation(self):
book = Book.objects.get(title="Test Book")
self.assertEqual(book.author, "Test Author")
def test_book_list_view(self):
response = self.client.get('/books/')
self.assertEqual(response.status_code, 200)
self.assertContains(response, "Test Book")
20. 测试可视化输出
测试生成图像或图表的代码:
- 比较哈希值或关键像素
- 检查生成文件的基本属性
- 使用专门的图像比较库
示例:
python复制import matplotlib.pyplot as plt
import unittest
import os
class TestPlot(unittest.TestCase):
def test_plot_creation(self):
plt.plot([1, 2, 3], [1, 4, 9])
output_path = 'test_plot.png'
plt.savefig(output_path)
self.assertTrue(os.path.exists(output_path))
self.assertGreater(os.path.getsize(output_path), 0)
# 清理
os.remove(output_path)
21. 测试多线程代码
测试并发代码的注意事项:
- 使用同步原语确保测试确定性
- 增加超时防止死锁
- 考虑使用unittest.mock模拟线程行为
示例:
python复制import threading
import unittest
import time
class TestThreading(unittest.TestCase):
def test_thread_execution(self):
result = []
def worker():
time.sleep(0.1)
result.append(1)
thread = threading.Thread(target=worker)
thread.start()
thread.join(timeout=1.0) # 设置超时
self.assertEqual(result, [1])
self.assertFalse(thread.is_alive())
22. 测试性能关键代码
虽然unittest不是性能测试框架,但可以用于基本性能验证:
python复制import timeit
import unittest
class TestPerformance(unittest.TestCase):
def test_fast_enough(self):
duration = timeit.timeit(
stmt='sorted(range(1000))',
number=1000
)
self.assertLess(duration, 1.0) # 1000次排序应小于1秒
23. 测试安全相关代码
安全测试的特殊考虑:
- 测试输入验证
- 测试敏感数据处理
- 测试权限控制
示例:
python复制import unittest
class TestSecurity(unittest.TestCase):
def test_password_hashing(self):
from hashlib import sha256
password = 'secret'
hashed = sha256(password.encode()).hexdigest()
self.assertNotEqual(hashed, password)
self.assertEqual(len(hashed), 64) # SHA-256哈希长度
def test_sql_injection(self):
user_input = "admin' --"
# 假设这是一个处理用户输入的SQL查询构建函数
query = build_query(user_input)
self.assertNotIn("--", query) # 应防止SQL注释注入
24. 测试配置管理
测试配置加载和验证:
python复制import unittest
import tempfile
import configparser
class TestConfig(unittest.TestCase):
def setUp(self):
self.config_file = tempfile.NamedTemporaryFile(mode='w+', delete=False)
self.config_file.write("""
[database]
host = localhost
port = 5432
""")
self.config_file.close()
def tearDown(self):
import os
os.unlink(self.config_file.name)
def test_config_loading(self):
config = configparser.ConfigParser()
config.read(self.config_file.name)
self.assertEqual(config['database']['host'], 'localhost')
self.assertEqual(config.getint('database', 'port'), 5432)
25. 测试API客户端
测试调用外部API的代码:
- 使用unittest.mock模拟requests
- 使用responses库
- 使用vcr.py记录和回放真实API调用
使用responses示例:
python复制import unittest
import responses
import requests
class TestAPIClient(unittest.TestCase):
@responses.activate
def test_api_call(self):
responses.add(
responses.GET,
'https://api.example.com/data',
json={'key': 'value'},
status=200
)
response = requests.get('https://api.example.com/data')
self.assertEqual(response.status_code, 200)
self.assertEqual(response.json(), {'key': 'value'})
26. 测试缓存行为
验证缓存相关代码:
python复制import unittest
from unittest.mock import patch
import time
class TestCache(unittest.TestCase):
@patch('cache_module.time.time', return_value=1000)
def test_cache_expiry(self, mock_time):
cache = Cache(expiry=60) # 60秒过期
cache.set('key', 'value')
# 模拟时间前进59秒
mock_time.return_value = 1059
self.assertEqual(cache.get('key'), 'value')
# 模拟时间前进61秒
mock_time.return_value = 1061
self.assertIsNone(cache.get('key'))
27. 测试序列化和反序列化
测试数据转换代码:
python复制import unittest
import json
class TestSerialization(unittest.TestCase):
def test_json_serialization(self):
data = {'key': 'value', 'num': 42}
json_str = json.dumps(data)
parsed = json.loads(json_str)
self.assertEqual(parsed['key'], 'value')
self.assertEqual(parsed['num'], 42)
self.assertIsInstance(parsed['num'], int)
def test_custom_serialization(self):
obj = CustomObject(value=10)
serialized = obj.serialize()
deserialized = CustomObject.deserialize(serialized)
self.assertEqual(deserialized.value, 10)
self.assertIsInstance(deserialized, CustomObject)
28. 测试算法实现
验证算法正确性:
python复制import unittest
class TestAlgorithms(unittest.TestCase):
def test_binary_search(self):
arr = [1, 3, 5, 7, 9]
self.assertEqual(binary_search(arr, 5), 2)
self.assertEqual(binary_search(arr, 1), 0)
self.assertEqual(binary_search(arr, 9), 4)
self.assertIsNone(binary_search(arr, 2))
def test_sort_stability(self):
items = [{'id': 2, 'name': 'b'}, {'id': 1, 'name': 'a'}, {'id': 2, 'name': 'c'}]
sorted_items = stable_sort(items, key=lambda x: x['id'])
# 检查排序稳定性:相同id保持原始顺序
self.assertEqual(sorted_items[1]['name'], 'b')
self.assertEqual(sorted_items[2]['name'], 'c')
29. 测试状态机
验证状态转换逻辑:
python复制import unittest
class TestStateMachine(unittest.TestCase):
def setUp(self):
self.sm = StateMachine(initial_state='idle')
def test_valid_transition(self):
self.sm.transition('start')
self.assertEqual(self.sm.state, 'running')
def test_invalid_transition(self):
with self.assertRaises(InvalidTransition):
self.sm.transition('stop') # 不能从idle直接到stopped
def test_guard_condition(self):
self.sm.set_ready(False)
with self.assertRaises(TransitionBlocked):
self.sm.transition('start')
30. 测试插件系统
验证插件架构:
python复制import unittest
from unittest.mock import MagicMock
class TestPluginSystem(unittest.TestCase):
def test_plugin_loading(self):
plugin_manager = PluginManager()
mock_plugin = MagicMock()
mock_plugin.name = 'test_plugin'
plugin_manager.register(mock_plugin)
plugin_manager.initialize_all()
mock_plugin.initialize.assert_called_once()
self.assertIn('test_plugin', plugin_manager.plugins)
def test_plugin_hook(self):
plugin_manager = PluginManager()
mock_plugin = MagicMock()
plugin_manager.register(mock_plugin)
plugin_manager.call_hook('pre_process', data={})
mock_plugin.pre_process.assert_called_once_with(data={})
31. 测试日志和监控
验证日志和监控集成:
python复制import unittest
from unittest.mock import patch
class TestMonitoring(unittest.TestCase):
@patch('monitoring.statsd.increment')
def test_counter_increment(self, mock_increment):
increment_counter('page.views')
mock_increment.assert_called_once_with('page.views')
@patch('logging.Logger.info')
def test_logging(self, mock_log):
log_important_event('system started')
mock_log.assert_called_once_with('system started')
32. 测试国际化(i18n)
验证多语言支持:
python复制import unittest
class TestInternationalization(unittest.TestCase):
def setUp(self):
self.translations = {
'en': {'greeting': 'Hello'},
'es': {'greeting': 'Hola'}
}
def test_english_translation(self):
self.assertEqual(
translate('greeting', 'en', self.translations),
'Hello'
)
def test_spanish_translation(self):
self.assertEqual(
translate('greeting', 'es', self.translations),
'Hola'
)
def test_missing_translation(self):
with self.assertRaises(TranslationMissing):
translate('nonexistent', 'en', self.translations)
33. 测试数据验证
验证输入数据处理:
python复制import unittest
class TestValidation(unittest.TestCase):
def test_email_validation(self):
self.assertTrue(is_valid_email('test@example.com'))
self.assertFalse(is_valid_email('invalid'))
self.assertFalse(is_valid_email('test@'))
self.assertFalse(is_valid_email('@example.com'))
def test_age_validation(self):
self.assertTrue(is_valid_age(25))
self.assertFalse(is_valid_age(-1))
self.assertFalse(is_valid_age(150))
self.assertTrue(is_valid_age(0)) # 新生儿
def test_password_strength(self):
self.assertEqual(check_password_strength(''), 'weak')
self.assertEqual(check_password_strength('password'), 'weak')
self.assertEqual(check_password_strength('Password1'), 'medium')
self.assertEqual(check_password_strength('P@ssw0rd!'), 'strong')
34. 测试日期计算
验证日期相关逻辑:
python复制import unittest
from datetime import date
class TestDateCalculations(unittest.TestCase):
def test_age_calculation(self):
today = date(2023, 1, 1)
birth_date = date(2000, 1, 1)
self.assertEqual(calculate_age(birth_date, today), 23)
birth_date = date(2000, 12, 31)
self.assertEqual(calculate_age(birth_date, today), 22)
def test_business_days(self):
# 周一至周五是工作日
self.assertEqual(count_business_days(
date(2023, 1, 2), # 周一
date(2023, 1, 6) # 周五
), 5)
# 包含周末
self.assertEqual(count_business_days(
date(2023, 1, 6), # 周五
date(2023, 1, 9) # 周一
), 1)
35. 测试文件格式处理
验证不同文件格式的解析:
python复制import unittest
import csv
import tempfile
class TestFileFormats(unittest.TestCase):
def test_csv_parsing(self):
csv_content = """name,age
Alice,30
Bob,25"""
with tempfile.NamedTemporaryFile(mode='w+') as f:
f.write(csv_content)
f.seek(0)
reader = csv.DictReader(f)
rows = list(reader)
self.assertEqual(len(rows), 2)
self.assertEqual(rows[0]['name'], 'Alice')
self.assertEqual(rows[0]['age'], '30')
self.assertEqual(rows[1]['name'], 'Bob')
self.assertEqual(rows[1]['age'], '25')
def test_json_parsing(self):
json_content = '{"key": "value", "num": 42}'
parsed = parse_json(json_content)
self.assertEqual(parsed['key'], 'value')
self.assertEqual(parsed['num'], 42)
self.assertIsInstance(parsed['num'], int)
36. 测试并发数据结构
验证线程安全数据结构:
python复制import unittest
import threading
class TestConcurrentQueue(unittest.TestCase):
def setUp(self):
self.queue = ConcurrentQueue()
def test_single_thread(self):
self.queue.put(1)
self.queue.put(2)
self.assertEqual(self.queue.get(), 1)
self.assertEqual(self.queue.get(), 2)
def test_multiple_threads(self):
def producer():
for i in range(100):
self.queue.put(i)
def consumer():
for _ in range(100):
item = self.queue.get()
self.assertIsNotNone(item)
threads = [
threading.Thread(target=producer),
threading.Thread(target=consumer)
]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertTrue(self.queue.empty())
37. 测试配置覆盖
验证配置覆盖逻辑:
python复制import unittest
class TestConfigOverride(unittest.TestCase):
def test_default_values(self):
config = load_config({})
self.assertEqual(config['timeout'], 30)
self.assertEqual(config['retries'], 3)
def test_override_values(self):
config = load_config({'timeout': 60, 'retries': 5})
self.assertEqual(config['timeout'], 60)
self.assertEqual(config['retries'], 5)
def test_invalid_values(self):
with self.assertRaises(ConfigError):
load_config({'timeout': -1})
with self.assertRaises(ConfigError):
load_config({'retries': 'invalid'})
38. 测试缓存失效
验证缓存失效策略:
python复制import unittest
import time
from unittest.mock import patch
class TestCacheInvalidation(unittest.TestCase):
@patch('time.time', return_value=1000)
def test_time_based_invalidation(self, mock_time):
cache = Cache(ttl=60)
cache.set('key', 'value')
# 在TTL内
mock_time.return_value = 1059
self.assertEqual(cache.get('key'), 'value')
# 超过TTL
mock_time.return_value = 1061
self.assertIsNone(cache.get('key'))
def test_event_based_invalidation(self):
cache = Cache()
cache.set('key', 'value')
cache.invalidate('key')
self.assertIsNone(cache.get('key'))
39. 测试错误恢复
验证错误处理和恢复机制:
python复制import unittest
from unittest.mock import patch, MagicMock
class TestErrorRecovery(unittest.TestCase):
@patch('database.connect')
def test_connection_retry(self, mock_connect):
# 第一次失败,第二次成功
mock_connect.side_effect = [ConnectionError, MagicMock()]
db = Database(max_retries=3)
connection = db.connect()
self.assertEqual(mock_connect.call_count, 2)
self.assertIsNotNone(connection)
def test_retry_exhausted(self):
with patch('database.connect', side_effect=ConnectionError) as mock_connect:
db = Database(max_retries=2)
with self.assertRaises(ConnectionError):
db.connect()
self.assertEqual(mock_connect.call_count, 2)
40. 测试数据迁移
验证数据迁移脚本:
python复制import unittest
import tempfile
import sqlite3
class TestDataMigration(unittest.TestCase):
def setUp(self):
# 创建源数据库
self.source_db = tempfile.NamedTemporaryFile(delete=False)
self.source_conn = sqlite3.connect(self.source_db.name)
self.source_cur = self.source_conn.cursor()
self.source_cur.execute("CREATE TABLE users (id INTEGER, name TEXT)")
self.source_cur.execute("INSERT INTO users VALUES (1, 'Alice')")
self.source_conn.commit()
