一、核心概念解析

1.1 基础定义:什么是单元测试?

单元测试是指对软件中的最小可测试单元(在Python中通常是函数或方法)进行检查和验证。unittest模块是Python自带的单元测试框架,它提供了:

  • 测试用例(TestCase):测试的基本单元,组织一组相关的测试方法
  • 测试固件(Fixture):测试前的准备和测试后的清理代码
  • 测试套件(TestSuite):测试用例的集合
  • 测试运行器(TestRunner):执行测试并展示结果

单元测试 vs 手动测试:

  • 手动测试:每次修改后都需要人工操作,容易遗漏,耗时耗力
  • 单元测试:自动化执行,快速反馈,确保每次修改都不会破坏现有功能
  • 手动测试:难以覆盖所有边界情况
  • 单元测试:可以轻松编写大量测试用例,覆盖各种场景

1.2 基本语法:你的第一个单元测试

让我们从一个简单的例子开始。假设我们有一个计算器类,我们要测试它的加法功能:

# calculator.py
class Calculator:
    """一个简单的计算器类"""
    
    def add(self, a, b):
        """返回两个数的和"""
        return a + b
    
    def subtract(self, a, b):
        """返回两个数的差"""
        return a - b
    
    def multiply(self, a, b):
        """返回两个数的积"""
        return a * b
    
    def divide(self, a, b):
        """返回两个数的商,b不能为0"""
        if b == 0:
            raise ValueError("除数不能为零")
        return a / b

现在,我们为这个计算器类编写单元测试:

# test_calculator.py
import unittest
from calculator import Calculator

class TestCalculator(unittest.TestCase):
    """测试Calculator类"""
    
    def test_add(self):
        """测试加法"""
        calc = Calculator()
        result = calc.add(3, 5)
        self.assertEqual(result, 8)  # 断言结果等于8
    
    def test_add_negative(self):
        """测试负数加法"""
        calc = Calculator()
        result = calc.add(-3, 5)
        self.assertEqual(result, 2)
    
    def test_add_float(self):
        """测试浮点数加法"""
        calc = Calculator()
        result = calc.add(3.5, 2.1)
        self.assertAlmostEqual(result, 5.6)  # 浮点数使用近似相等断言
    
    def test_subtract(self):
        """测试减法"""
        calc = Calculator()
        result = calc.subtract(10, 4)
        self.assertEqual(result, 6)
    
    def test_multiply(self):
        """测试乘法"""
        calc = Calculator()
        result = calc.multiply(3, 7)
        self.assertEqual(result, 21)
    
    def test_divide(self):
        """测试除法"""
        calc = Calculator()
        result = calc.divide(10, 2)
        self.assertEqual(result, 5)
    
    def test_divide_by_zero(self):
        """测试除零异常"""
        calc = Calculator()
        with self.assertRaises(ValueError):  # 断言会抛出ValueError异常
            calc.divide(10, 0)

if __name__ == '__main__':
    unittest.main()

运行测试:

python test_calculator.py

你会看到类似这样的输出:

.......
----------------------------------------------------------------------
Ran 7 tests in 0.001s

OK

每个点表示一个通过的测试。如果有测试失败,会显示F和详细的失败信息。

1.3 核心特点:为什么要使用unittest?

  1. 自动化测试:一次编写,多次运行,节省手动测试时间
  2. 快速反馈:修改代码后立即运行测试,快速发现错误
  3. 文档作用:测试用例本身就是如何使用代码的示例
  4. 设计促进:编写测试会促使你思考代码的接口和边界情况
  5. 重构保障:有了测试套件,重构代码时更有信心
  6. 团队协作:确保团队成员的修改不会破坏现有功能

二、应用场景详解

2.1 测试固件:setUp和tearDown

在测试中,我们经常需要一些公共的准备和清理代码。unittest提供了setUp和tearDown方法来解决这个问题。

  • setUp():在每个测试方法之前运行,用于准备测试环境
  • tearDown():在每个测试方法之后运行,用于清理测试环境
  • setUpClass():在整个测试类开始前运行一次(类方法)
  • tearDownClass():在整个测试类结束后运行一次(类方法)
# test_with_fixtures.py
import unittest
import tempfile
import os

class FileProcessor:
    """一个处理文件的类"""
    
    def process_file(self, filepath):
        """处理文件,返回文件内容"""
        with open(filepath, 'r') as f:
            content = f.read()
        return content.upper()  # 转换为大写
    
    def count_lines(self, filepath):
        """计算文件行数"""
        with open(filepath, 'r') as f:
            lines = f.readlines()
        return len(lines)

class TestFileProcessor(unittest.TestCase):
    """测试FileProcessor类,演示测试固件的使用"""
    
    @classmethod
    def setUpClass(cls):
        """整个测试类开始前执行一次"""
        print("=== 开始FileProcessor测试 ===")
        cls.temp_dir = tempfile.mkdtemp()  # 创建临时目录
        print(f"临时目录: {cls.temp_dir}")
    
    @classmethod
    def tearDownClass(cls):
        """整个测试类结束后执行一次"""
        # 清理临时目录
        import shutil
        shutil.rmtree(cls.temp_dir)
        print(f"清理临时目录: {cls.temp_dir}")
        print("=== 结束FileProcessor测试 ===")
    
    def setUp(self):
        """每个测试方法前执行"""
        # 创建测试文件
        self.test_file = os.path.join(self.temp_dir, 'test.txt')
        with open(self.test_file, 'w') as f:
            f.write("Hello, World!\nThis is a test.\nHave a nice day!")
        
        # 创建FileProcessor实例
        self.processor = FileProcessor()
        
        # 每个测试方法都会收到一个全新的processor实例
        print(f"\n设置测试环境: {self.test_file}")
    
    def tearDown(self):
        """每个测试方法后执行"""
        # 删除测试文件
        if os.path.exists(self.test_file):
            os.remove(self.test_file)
        print(f"清理测试文件: {self.test_file}")
    
    def test_process_file(self):
        """测试文件处理"""
        result = self.processor.process_file(self.test_file)
        expected = "HELLO, WORLD!\nTHIS IS A TEST.\nHAVE A NICE DAY!"
        self.assertEqual(result, expected)
    
    def test_count_lines(self):
        """测试行数统计"""
        line_count = self.processor.count_lines(self.test_file)
        self.assertEqual(line_count, 3)
    
    def test_empty_file(self):
        """测试空文件"""
        # 创建一个空文件
        empty_file = os.path.join(self.temp_dir, 'empty.txt')
        with open(empty_file, 'w') as f:
            f.write("")
        
        try:
            # 测试空文件处理
            result = self.processor.process_file(empty_file)
            self.assertEqual(result, "")
            
            # 测试空文件行数
            line_count = self.processor.count_lines(empty_file)
            self.assertEqual(line_count, 0)
        finally:
            # 清理空文件
            if os.path.exists(empty_file):
                os.remove(empty_file)

if __name__ == '__main__':
    unittest.main(verbosity=2)  # 显示更详细的输出

2.2 断言方法:验证测试结果

断言是测试的核心,unittest提供了丰富的断言方法:

# test_assertions.py
import unittest

class TestAssertions(unittest.TestCase):
    """演示各种断言方法的使用"""
    
    def test_equality(self):
        """相等性断言"""
        self.assertEqual(3 + 4, 7)  # 是否相等
        self.assertNotEqual(3 + 4, 8)  # 是否不相等
    
    def test_truthiness(self):
        """真值断言"""
        self.assertTrue(1 < 2)  # 是否为真
        self.assertFalse(1 > 2)  # 是否为假
    
    def test_none(self):
        """None值断言"""
        value = None
        self.assertIsNone(value)  # 是否为None
        
        value = 42
        self.assertIsNotNone(value)  # 是否不为None
    
    def test_identity(self):
        """同一性断言(is/is not)"""
        a = [1, 2, 3]
        b = a
        c = [1, 2, 3]
        
        self.assertIs(a, b)  # 是否是同一个对象
        self.assertIsNot(a, c)  # 是否不是同一个对象
    
    def test_containment(self):
        """包含关系断言"""
        my_list = [1, 2, 3, 4, 5]
        my_dict = {'a': 1, 'b': 2}
        
        self.assertIn(3, my_list)  # 是否包含
        self.assertNotIn(6, my_list)  # 是否不包含
        self.assertIn('a', my_dict)  # 字典中是否包含键
    
    def test_comparison(self):
        """比较断言"""
        self.assertGreater(5, 3)  # 是否大于
        self.assertGreaterEqual(5, 5)  # 是否大于等于
        self.assertLess(3, 5)  # 是否小于
        self.assertLessEqual(3, 3)  # 是否小于等于
    
    def test_exceptions(self):
        """异常断言"""
        # 断言会抛出特定异常
        with self.assertRaises(ZeroDivisionError):
            result = 1 / 0
        
        # 断言异常消息包含特定文本
        with self.assertRaises(ValueError) as context:
            int('invalid')
        
        self.assertIn('invalid literal', str(context.exception))
    
    def test_approximate(self):
        """近似相等断言(用于浮点数)"""
        result = 0.1 + 0.2
        self.assertAlmostEqual(result, 0.3)  # 默认精度7位小数
        self.assertAlmostEqual(result, 0.3, places=15)  # 指定精度
    
    def test_type(self):
        """类型断言"""
        value = 42
        self.assertIsInstance(value, int)  # 是否是特定类型
        self.assertNotIsInstance(value, str)  # 是否不是特定类型
    
    def test_collections(self):
        """集合断言"""
        expected = [1, 2, 3, 4, 5]
        actual = [1, 2, 3, 4, 5]
        
        self.assertListEqual(expected, actual)  # 列表是否相等
        self.assertSequenceEqual(expected, actual)  # 序列是否相等
        
        expected_dict = {'a': 1, 'b': 2}
        actual_dict = {'b': 2, 'a': 1}  # 顺序不同但相等
        self.assertDictEqual(expected_dict, actual_dict)  # 字典是否相等
        
        expected_set = {1, 2, 3}
        actual_set = {3, 2, 1}  # 集合顺序无关
        self.assertSetEqual(expected_set, actual_set)  # 集合是否相等

if __name__ == '__main__':
    unittest.main(verbosity=2)

2.3 跳过测试和预期失败

有时我们需要临时跳过某些测试,或者标记某些测试为预期失败:

# test_skip_expected_failure.py
import unittest
import sys

class TestSkipAndExpectedFailure(unittest.TestCase):
    """演示跳过测试和预期失败"""
    
    def test_normal(self):
        """正常测试"""
        self.assertEqual(1 + 1, 2)
    
    @unittest.skip("跳过这个测试,因为功能尚未实现")
    def test_skip_with_reason(self):
        """被跳过的测试"""
        self.fail("这个测试不应该运行")
    
    @unittest.skipIf(sys.version_info < (3, 7), "需要Python 3.7或更高版本")
    def test_skip_conditionally(self):
        """条件跳过测试"""
        # 这个测试只在Python 3.7+上运行
        self.assertTrue(sys.version_info >= (3, 7))
    
    @unittest.skipUnless(sys.platform.startswith("win"), "需要Windows系统")
    def test_windows_only(self):
        """Windows专用测试"""
        # 这个测试只在Windows上运行
        import platform
        self.assertTrue(platform.system() == "Windows")
    
    @unittest.expectedFailure
    def test_expected_failure(self):
        """预期失败的测试"""
        # 这个测试目前会失败,但我们知道并标记为预期失败
        # 当它通过时,会被报告为意外通过
        self.assertEqual(1 + 1, 3)  # 这明显是错的
    
    def test_may_skip(self):
        """测试中动态跳过"""
        # 某些条件下跳过
        if not hasattr(sys, 'getwindowsversion'):
            self.skipTest("这个测试需要Windows特定功能")
        
        # 如果满足条件,继续测试
        import platform
        self.assertEqual(platform.system(), "Windows")

if __name__ == '__main__':
    unittest.main(verbosity=2)

三、高级技巧

3.1 测试发现和测试套件

unittest可以自动发现和运行测试:

# 自动发现并运行当前目录下所有test_*.py文件中的测试
python -m unittest discover

# 运行特定模块的测试
python -m unittest test_calculator

# 运行特定测试类
python -m unittest test_calculator.TestCalculator

# 运行特定测试方法
python -m unittest test_calculator.TestCalculator.test_add

# 运行指定目录下的测试
python -m unittest discover -s tests

# 使用模式匹配发现测试文件
python -m unittest discover -p "*_test.py"

# 显示详细输出
python -m unittest discover -v

你也可以编程方式创建测试套件:

# test_suite_demo.py
import unittest
from test_calculator import TestCalculator
from test_assertions import TestAssertions

def create_test_suite():
    """创建自定义测试套件"""
    suite = unittest.TestSuite()
    
    # 添加单个测试方法
    suite.addTest(TestCalculator('test_add'))
    
    # 添加整个测试类
    suite.addTest(unittest.makeSuite(TestCalculator))
    
    # 添加多个测试类
    suite.addTests([
        unittest.makeSuite(TestCalculator),
        unittest.TestLoader().loadTestsFromTestCase(TestAssertions)
    ])
    
    return suite

if __name__ == '__main__':
    # 创建测试运行器
    runner = unittest.TextTestRunner(verbosity=2)
    
    # 运行自定义测试套件
    suite = create_test_suite()
    result = runner.run(suite)
    
    # 输出测试结果统计
    print(f"\n运行了 {result.testsRun} 个测试")
    print(f"失败: {len(result.failures)} 个")
    print(f"错误: {len(result.errors)} 个")
    print(f"跳过: {len(result.skipped)} 个")

3.2 Mock和补丁:隔离测试对象

单元测试应该专注于测试单个单元,而不是它的依赖。unittest.mock模块提供了Mock对象和补丁功能,用于隔离测试代码:

# test_with_mock.py
import unittest
from unittest.mock import Mock, MagicMock, patch, call
from datetime import datetime

class WeatherService:
    """模拟的天气服务"""
    
    def get_temperature(self, city):
        """获取城市温度(这里模拟为总是返回20)"""
        # 实际中这里可能会有网络请求
        return 20

class WeatherReporter:
    """天气报告器,依赖WeatherService"""
    
    def __init__(self, weather_service):
        self.weather_service = weather_service
    
    def report(self, city):
        """生成天气报告"""
        temp = self.weather_service.get_temperature(city)
        return f"{city}的当前温度是{temp}°C"

class TestWeatherReporter(unittest.TestCase):
    """测试WeatherReporter,使用Mock隔离WeatherService"""
    
    def test_report_with_mock(self):
        """使用Mock对象测试"""
        # 创建Mock对象替代真实的WeatherService
        mock_service = Mock(spec=WeatherService)
        mock_service.get_temperature.return_value = 25
        
        # 创建被测试对象,注入Mock
        reporter = WeatherReporter(mock_service)
        
        # 调用被测试方法
        result = reporter.report("北京")
        
        # 验证结果
        self.assertEqual(result, "北京的当前温度是25°C")
        
        # 验证Mock被正确调用
        mock_service.get_temperature.assert_called_once_with("北京")
    
    def test_report_with_magic_mock(self):
        """使用MagicMock测试"""
        # MagicMock是Mock的增强版,支持魔术方法
        mock_service = MagicMock()
        mock_service.get_temperature.return_value = 30
        
        reporter = WeatherReporter(mock_service)
        result = reporter.report("上海")
        
        self.assertEqual(result, "上海的当前温度是30°C")
        mock_service.get_temperature.assert_called_once_with("上海")
    
    @patch('__main__.WeatherService')  # 补丁WeatherService类
    def test_report_with_patch_decorator(self, MockWeatherService):
        """使用patch装饰器测试"""
        # 创建Mock实例
        mock_instance = MockWeatherService.return_value
        mock_instance.get_temperature.return_value = 18
        
        # 创建被测试对象
        reporter = WeatherReporter(mock_instance)
        result = reporter.report("广州")
        
        # 验证
        self.assertEqual(result, "广州的当前温度是18°C")
        mock_instance.get_temperature.assert_called_once_with("广州")
    
    def test_report_with_patch_context(self):
        """使用patch上下文管理器测试"""
        with patch('__main__.WeatherService') as MockWeatherService:
            # 在上下文管理器内部,WeatherService被替换为Mock
            mock_instance = MockWeatherService.return_value
            mock_instance.get_temperature.return_value = 22
            
            reporter = WeatherReporter(mock_instance)
            result = reporter.report("深圳")
            
            self.assertEqual(result, "深圳的当前温度是22°C")
    
    def test_multiple_calls(self):
        """测试多次调用"""
        mock_service = Mock()
        
        # 设置多次调用的不同返回值
        mock_service.get_temperature.side_effect = [20, 25, 30]
        
        reporter = WeatherReporter(mock_service)
        
        # 第一次调用
        result1 = reporter.report("北京")
        self.assertEqual(result1, "北京的当前温度是20°C")
        
        # 第二次调用
        result2 = reporter.report("上海")
        self.assertEqual(result2, "上海的当前温度是25°C")
        
        # 第三次调用
        result3 = reporter.report("广州")
        self.assertEqual(result3, "广州的当前温度是30°C")
        
        # 验证调用次数和参数
        expected_calls = [call("北京"), call("上海"), call("广州")]
        mock_service.get_temperature.assert_has_calls(expected_calls)
        self.assertEqual(mock_service.get_temperature.call_count, 3)
    
    def test_exception(self):
        """测试异常情况"""
        mock_service = Mock()
        
        # 设置抛出异常
        mock_service.get_temperature.side_effect = Exception("网络错误")
        
        reporter = WeatherReporter(mock_service)
        
        # 验证异常被传播
        with self.assertRaises(Exception) as context:
            reporter.report("北京")
        
        self.assertIn("网络错误", str(context.exception))

if __name__ == '__main__':
    unittest.main(verbosity=2)

3.3 测试驱动开发(TDD)示例

测试驱动开发(TDD)是一种先写测试,再写实现代码的开发方法:

# tdd_example.py
import unittest

# 我们先写测试,再写实现

class TestFizzBuzz(unittest.TestCase):
    """FizzBuzz游戏的测试"""
    
    def test_fizzbuzz(self):
        """测试FizzBuzz游戏逻辑"""
        fb = FizzBuzz()
        
        # 测试普通数字
        self.assertEqual(fb.convert(1), "1")
        self.assertEqual(fb.convert(2), "2")
        
        # 测试3的倍数
        self.assertEqual(fb.convert(3), "Fizz")
        self.assertEqual(fb.convert(6), "Fizz")
        
        # 测试5的倍数
        self.assertEqual(fb.convert(5), "Buzz")
        self.assertEqual(fb.convert(10), "Buzz")
        
        # 测试3和5的公倍数
        self.assertEqual(fb.convert(15), "FizzBuzz")
        self.assertEqual(fb.convert(30), "FizzBuzz")
        
        # 测试边界情况
        with self.assertRaises(ValueError):
            fb.convert(0)
        
        with self.assertRaises(ValueError):
            fb.convert(-1)
        
        with self.assertRaises(TypeError):
            fb.convert("不是数字")

# 运行测试会失败,因为FizzBuzz类还不存在
# 现在我们来实现它

class FizzBuzz:
    """FizzBuzz游戏"""
    
    def convert(self, n):
        """
        将数字转换为FizzBuzz字符串:
        - 3的倍数 -> "Fizz"
        - 5的倍数 -> "Buzz"
        - 3和5的公倍数 -> "FizzBuzz"
        - 其他 -> 数字字符串
        """
        # 参数检查
        if not isinstance(n, int):
            raise TypeError("输入必须是整数")
        
        if n < 1:
            raise ValueError("输入必须是正整数")
        
        # FizzBuzz逻辑
        result = ""
        
        if n % 3 == 0:
            result += "Fizz"
        
        if n % 5 == 0:
            result += "Buzz"
        
        return result if result else str(n)

# 现在再次运行测试应该通过
if __name__ == '__main__':
    unittest.main(verbosity=2)

四、实战案例:测试一个简单的购物车

让我们看一个完整的实战例子,测试一个购物车系统:

# shopping_cart.py
class ShoppingCart:
    """简单的购物车类"""
    
    def __init__(self):
        self.items = {}  # 商品名 -> 数量
        self.prices = {  # 商品价格
            'apple': 5.0,
            'banana': 3.0,
            'orange': 4.0,
            'milk': 10.0
        }
    
    def add_item(self, item, quantity=1):
        """添加商品到购物车"""
        if item not in self.prices:
            raise ValueError(f"商品 '{item}' 不存在")
        
        if quantity <= 0:
            raise ValueError("数量必须大于0")
        
        if item in self.items:
            self.items[item] += quantity
        else:
            self.items[item] = quantity
    
    def remove_item(self, item, quantity=1):
        """从购物车移除商品"""
        if item not in self.items:
            raise ValueError(f"购物车中没有商品 '{item}'")
        
        if quantity <= 0:
            raise ValueError("数量必须大于0")
        
        if quantity >= self.items[item]:
            del self.items[item]
        else:
            self.items[item] -= quantity
    
    def get_total(self):
        """计算购物车总价"""
        total = 0.0
        for item, quantity in self.items.items():
            if item in self.prices:
                total += self.prices[item] * quantity
        return round(total, 2)
    
    def get_item_count(self, item=None):
        """获取商品数量"""
        if item is None:
            return sum(self.items.values())
        elif item in self.items:
            return self.items[item]
        else:
            return 0
    
    def clear(self):
        """清空购物车"""
        self.items.clear()
    
    def apply_discount(self, percentage):
        """应用折扣"""
        if not 0 <= percentage <= 100:
            raise ValueError("折扣百分比必须在0-100之间")
        
        total = self.get_total()
        discount = total * (percentage / 100)
        return round(total - discount, 2)

现在为这个购物车编写测试:

# test_shopping_cart.py
import unittest
from shopping_cart import ShoppingCart

class TestShoppingCart(unittest.TestCase):
    """测试ShoppingCart类"""
    
    def setUp(self):
        """每个测试前的准备"""
        self.cart = ShoppingCart()
        print("创建新的购物车实例")
    
    def tearDown(self):
        """每个测试后的清理"""
        self.cart.clear()
        print("清空购物车")
    
    def test_add_item(self):
        """测试添加商品"""
        self.cart.add_item('apple', 2)
        self.assertEqual(self.cart.get_item_count('apple'), 2)
        self.assertEqual(self.cart.get_item_count(), 2)
    
    def test_add_multiple_items(self):
        """测试添加多种商品"""
        self.cart.add_item('apple', 2)
        self.cart.add_item('banana', 3)
        
        self.assertEqual(self.cart.get_item_count('apple'), 2)
        self.assertEqual(self.cart.get_item_count('banana'), 3)
        self.assertEqual(self.cart.get_item_count(), 5)
    
    def test_add_existing_item(self):
        """测试添加已存在的商品"""
        self.cart.add_item('apple', 2)
        self.cart.add_item('apple', 3)  # 再次添加
        
        self.assertEqual(self.cart.get_item_count('apple'), 5)
    
    def test_remove_item(self):
        """测试移除商品"""
        self.cart.add_item('apple', 5)
        self.cart.remove_item('apple', 2)
        
        self.assertEqual(self.cart.get_item_count('apple'), 3)
    
    def test_remove_item_completely(self):
        """测试完全移除商品"""
        self.cart.add_item('apple', 3)
        self.cart.remove_item('apple', 5)  # 移除数量超过现有数量
        
        self.assertEqual(self.cart.get_item_count('apple'), 0)
        self.assertNotIn('apple', self.cart.items)
    
    def test_remove_nonexistent_item(self):
        """测试移除不存在的商品"""
        with self.assertRaises(ValueError) as context:
            self.cart.remove_item('nonexistent')
        
        self.assertIn("购物车中没有商品", str(context.exception))
    
    def test_add_invalid_item(self):
        """测试添加无效商品"""
        with self.assertRaises(ValueError) as context:
            self.cart.add_item('nonexistent')
        
        self.assertIn("商品 'nonexistent' 不存在", str(context.exception))
    
    def test_add_invalid_quantity(self):
        """测试添加无效数量"""
        with self.assertRaises(ValueError) as context:
            self.cart.add_item('apple', 0)  # 数量为0
        
        self.assertIn("数量必须大于0", str(context.exception))
        
        with self.assertRaises(ValueError) as context:
            self.cart.add_item('apple', -1)  # 数量为负
        
        self.assertIn("数量必须大于0", str(context.exception))
    
    def test_get_total(self):
        """测试计算总价"""
        self.cart.add_item('apple', 2)  # 2 * 5.0 = 10.0
        self.cart.add_item('banana', 3)  # 3 * 3.0 = 9.0
        self.cart.add_item('milk', 1)   # 1 * 10.0 = 10.0
        
        total = self.cart.get_total()
        expected = 10.0 + 9.0 + 10.0  # 29.0
        
        self.assertEqual(total, expected)
        self.assertAlmostEqual(total, 29.0)  # 使用浮点数断言
    
    def test_get_total_empty_cart(self):
        """测试空购物车的总价"""
        total = self.cart.get_total()
        self.assertEqual(total, 0.0)
    
    def test_clear_cart(self):
        """测试清空购物车"""
        self.cart.add_item('apple', 2)
        self.cart.add_item('banana', 3)
        
        self.assertEqual(self.cart.get_item_count(), 5)
        
        self.cart.clear()
        
        self.assertEqual(self.cart.get_item_count(), 0)
        self.assertEqual(self.cart.get_total(), 0.0)
    
    def test_apply_discount(self):
        """测试应用折扣"""
        self.cart.add_item('apple', 4)  # 4 * 5.0 = 20.0
        self.cart.add_item('banana', 2)  # 2 * 3.0 = 6.0
        
        # 总价26.0,打8折
        discounted = self.cart.apply_discount(20)  # 20%折扣
        
        expected = 26.0 * 0.8  # 20.8
        self.assertAlmostEqual(discounted, expected)
    
    def test_apply_invalid_discount(self):
        """测试应用无效折扣"""
        with self.assertRaises(ValueError) as context:
            self.cart.apply_discount(-10)  # 负折扣
        
        self.assertIn("折扣百分比必须在0-100之间", str(context.exception))
        
        with self.assertRaises(ValueError) as context:
            self.cart.apply_discount(150)  # 超过100%的折扣
        
        self.assertIn("折扣百分比必须在0-100之间", str(context.exception))
    
    def test_edge_cases(self):
        """测试边界情况"""
        # 测试大量商品
        for i in range(100):
            self.cart.add_item('apple', 1)
        
        self.assertEqual(self.cart.get_item_count('apple'), 100)
        self.assertEqual(self.cart.get_total(), 500.0)  # 100 * 5.0
        
        # 清空后重新添加
        self.cart.clear()
        self.assertEqual(self.cart.get_item_count(), 0)
        
        # 测试只添加一个商品
        self.cart.add_item('orange')
        self.assertEqual(self.cart.get_item_count('orange'), 1)
        
        # 移除这个商品
        self.cart.remove_item('orange')
        self.assertEqual(self.cart.get_item_count('orange'), 0)

if __name__ == '__main__':
    # 使用TestLoader来组织测试
    loader = unittest.TestLoader()
    
    # 创建测试套件
    suite = loader.loadTestsFromTestCase(TestShoppingCart)
    
    # 创建测试运行器
    runner = unittest.TextTestRunner(verbosity=2)
    
    # 运行测试
    print("开始运行购物车测试...")
    print("=" * 50)
    result = runner.run(suite)
    print("=" * 50)
    
    # 输出测试结果总结
    if result.wasSuccessful():
        print("✓ 所有测试通过!")
    else:
        print(f"✗ 测试失败: {len(result.failures)} 个失败, {len(result.errors)} 个错误")

五、注意事项

5.1 使用限制

  1. 测试覆盖:unittest本身不提供测试覆盖率统计,需要配合coverage等工具
  2. 测试隔离:测试之间应该完全独立,不共享状态
  3. 测试速度:测试应该快速运行,避免I/O操作和网络请求
  4. 测试数据:使用测试固件或Mock来准备测试数据,避免依赖外部系统
  5. 测试命名:测试方法名应该清晰描述测试内容,通常以test_开头

5.2 常见问题

Q: 测试文件应该放在哪里?

A: 通常放在项目根目录的tests目录中,或者与源代码放在同一目录,以test_开头命名。

Q: 如何组织大型测试套件?

A: 使用测试目录结构,unittest的discover命令可以自动发现测试。

Q: 测试应该有多详细?

A: 测试应该覆盖正常情况、边界情况和异常情况,但也要保持简洁。

Q: 如何处理测试中的随机性?

A: 使用随机种子,或者使用Mock固定随机结果。

Q: 测试数据库操作怎么办?

A: 使用内存数据库(如SQLite内存数据库),或者使用Mock模拟数据库操作。

Q: 什么时候写测试?

A: 理想情况下,在写实现代码之前写测试(TDD),但至少应该在代码完成后立即写测试。

5.3 替代方案

  1. pytest:更简洁的语法,丰富的插件生态系统
  2. nose2:unittest的扩展,提供更多功能
  3. doctest:将文档字符串中的示例作为测试
  4. hypothesis:基于属性的测试,自动生成测试用例

何时选择替代方案:

  • 需要更简洁的语法 → 使用pytest
  • 需要参数化测试、固件依赖等高级功能 → 使用pytest
  • 需要基于属性的测试 → 使用hypothesis
  • 项目已使用unittest,需要向后兼容 → 继续使用unittest

六、总结

单元测试是保证代码质量的关键实践。通过本文的学习,你应该已经掌握了:

  • ✅ unittest基础:如何编写测试类、测试方法和断言
  • ✅ 测试固件:使用setUp、tearDown管理测试环境
  • ✅ 测试组织:测试发现、测试套件和测试运行
  • ✅ Mock和补丁:隔离测试对象,专注测试单元本身
  • ✅ 测试驱动开发:先写测试,再写实现代码
  • ✅ 实战应用:为实际项目编写全面的测试

测试心法:

  1. 测试要快:快速运行,快速反馈
  2. 测试要独立:测试之间不相互依赖
  3. 测试要可重复:每次运行结果应该一致
  4. 测试要全面:覆盖正常、边界和异常情况
  5. 测试即文档:测试应该清晰地说明代码的预期行为

你在编写单元测试时遇到过哪些挑战?有什么好的测试实践想分享?欢迎在评论区交流你的测试经验和心得!

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐