Skip to content

第21章 测试策略

pybind11 项目的测试需要同时覆盖 Python 接口和 C++ 实现。良好的测试策略能在早期发现绑定错误,避免生产环境中的问题。

pybind11 项目本身使用 Catch2 作为测试框架,测试代码位于 pybind11/tests/ 目录中。

tests/test_basic_binding.cpp
#include <pybind11/pybind11.h>
#include <catch2/catch.hpp>
namespace py = pybind11;
// 被测函数
int add(int a, int b) {
return a + b;
}
int multiply(int a, int b) {
return a * b;
}
// C++ 单元测试
TEST_CASE("add function") {
REQUIRE(add(2, 3) == 5);
REQUIRE(add(0, 0) == 0);
REQUIRE(add(-1, 1) == 0);
}
TEST_CASE("multiply function") {
REQUIRE(multiply(2, 3) == 6);
REQUIRE(multiply(0, 100) == 0);
REQUIRE(multiply(-2, 3) == -6);
}
// 绑定代码
PYBIND11_MODULE(test_basic_module, m) {
m.def("add", &add);
m.def("multiply", &multiply);
}
import pytest
import test_basic_module as m
class TestAdd:
def test_positive(self):
assert m.add(2, 3) == 5
def test_zero(self):
assert m.add(0, 0) == 0
def test_negative(self):
assert m.add(-1, 1) == 0
class TestMultiply:
def test_positive(self):
assert m.multiply(2, 3) == 6
def test_zero(self):
assert m.multiply(0, 100) == 0
def test_negative(self):
assert m.multiply(-2, 3) == -6

关键洞察:C++ 端测试可以验证核心逻辑,Python 端测试可以验证绑定正确性。两者结合才能确保绑定的完整性。

pytest 是 Python 测试的事实标准,与 pybind11 模块配合良好。

Terminal window
pip install pytest-cpp
pytest tests/ --cpp-glob="*.cpp" --cpp-tests
[tool.pytest.ini_options]
testpaths = ["tests"]
python_files = "test_*.py"
cpp_files = ["tests/*.cpp"]
python_files = "test_*.py"
addopts = "-v --tb=short"
import pytest
import my_module as m
@pytest.mark.parametrize("a, b, expected", [
(1, 2, 3),
(0, 0, 0),
(-1, 1, 0),
(100, 200, 300),
])
def test_add(a, b, expected):
assert m.add(a, b) == expected
@pytest.mark.parametrize("input, output", [
([1, 2, 3], 6),
([], 0),
([-1, -2, -3], -6),
([0.5, 0.5], 1.0),
])
def test_sum(input, output):
assert abs(m.sum(input) - output) < 1e-10
import pytest
import my_module
@pytest.fixture(scope="module")
def module():
"""模块级 fixture,用于所有测试共享模块"""
return my_module
@pytest.fixture
def vector_data():
"""生成测试用向量数据"""
return [1.0, 2.0, 3.0, 4.0, 5.0]
class TestVectorOperations:
def test_sum(self, module, vector_data):
result = module.sum(vector_data)
assert abs(result - 15.0) < 1e-10
def test_mean(self, module, vector_data):
result = module.mean(vector_data)
assert abs(result - 3.0) < 1e-10

关键洞察:使用 pytest fixtures 管理 C++ 模块的加载和测试数据,可以减少重复代码并提高测试速度。模块级 fixture 可以避免重复导入开销。

21.3 C++ 单元测试(gtest/gbench 集成)

Section titled “21.3 C++ 单元测试(gtest/gbench 集成)”

Google Test (gtest) 是 C++ 测试的标准框架,可以与 pybind11 一起使用。

tests/test_with_gtest.cpp
#include <pybind11/pybind11.h>
#include <gtest/gtest.h>
namespace py = pybind11;
// 被测类
class Calculator {
public:
double add(double a, double b) const { return a + b; }
double divide(double a, double b) {
if (b == 0.0) throw std::runtime_error("division by zero");
return a / b;
}
};
// gtest 测试用例
TEST(CalculatorTest, Add) {
Calculator calc;
EXPECT_DOUBLE_EQ(calc.add(1.0, 2.0), 3.0);
EXPECT_DOUBLE_EQ(calc.add(-1.0, 1.0), 0.0);
}
TEST(CalculatorTest, Divide) {
Calculator calc;
EXPECT_DOUBLE_EQ(calc.divide(10.0, 2.0), 5.0);
}
TEST(CalculatorTest, DivideByZero) {
Calculator calc;
EXPECT_THROW(calc.divide(1.0, 0.0), std::runtime_error);
}
// 绑定
PYBIND11_MODULE(gtest_module, m) {
py::class_<Calculator>(m, "Calculator")
.def(py::init<>())
.def("add", &Calculator::add)
.def("divide", &Calculator::divide);
// 注册 gtest(让 Python 可以调用)
py::exec(R"(
import sys
try:
import gtest
except ImportError:
pass
)", py::globals());
}
tests/test_with_benchmark.cpp
#include <pybind11/pybind11.h>
#include <benchmark/benchmark.h>
namespace py = pybind11;
// 被测函数
double compute_expensive(double x) {
double result = 0.0;
for (int i = 0; i < 1000; ++i) {
result += std::sin(x * i);
}
return result;
}
// benchmark 测试
static void BM_Compute(benchmark::State& state) {
double x = 1.5;
for (auto _ : state) {
benchmark::DoNotOptimize(compute_expensive(x));
}
}
BENCHMARK(BM_Compute);
BENCHMARK_MAIN();
import subprocess
import sys
def test_cpp_gtest():
"""运行 C++ gtest 测试"""
result = subprocess.run(
[sys.executable, "-m", "gtest", "--verbose"],
capture_output=True,
text=True
)
assert result.returncode == 0, f"gtest failed: {result.stderr}"
def test_cpp_benchmark():
"""运行 C++ benchmark"""
result = subprocess.run(
["./benchmark_binary", "--benchmark_format=json"],
capture_output=True,
text=True
)
assert result.returncode == 0

关键洞察:gtest 提供强大的 C++ 单元测试能力,与 pytest 可以很好地配合。对于性能关键代码,使用 Google Benchmark 进行微基准测试。

使用 coverage.py 和 gcov 测量测试覆盖率。

Terminal window
pip install coverage
coverage run -m pytest tests/
coverage report
coverage html
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} --coverage")
set(CMAKE_EXE_LINKER_FLAGS "${CMAKE_EXE_LINKER_FLAGS} --coverage")
pybind11_add_module(my_module my_module.cpp)
target_link_libraries(my_module PRIVATE gcov)
Terminal window
mkdir build && cd build
cmake -DCMAKE_BUILD_TYPE=Debug ..
make
python3 -m pytest ../tests/
gcov ../my_module.cpp -o .
lcov --capture --directory . --output-file coverage.info
genhtml coverage.info --output-directory html_coverage
Terminal window
coverage run --append -m pytest tests/
coverage combine
coverage report

关键洞察:覆盖率报告帮助识别测试盲区。目标是高覆盖率,但更要关注边界条件和错误路径的覆盖。

持续集成确保每次提交都不会破坏已有功能。

name: Test
on:
push:
branches: [main]
pull_request:
branches: [main]
jobs:
test:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v3
- name: Setup Python
uses: actions/setup-python@v4
with:
python-version: ['3.8', '3.9', '3.10', '3.11']
- name: Install dependencies
run: |
pip install pytest pytest-cov pybind11 numpy
- name: Build C++ extension
run: |
mkdir build && cd build
cmake ..
make
- name: Run Python tests
run: |
pytest tests/ -v --cov=my_module --cov-report=xml
- name: Run C++ tests
run: |
./build/tests/cpptest
- name: Upload coverage
uses: codecov/codecov-action@v3
with:
file: ./coverage.xml
jobs:
build:
strategy:
matrix:
os: [ubuntu-latest, macos-latest, windows-latest]
python-version: ['3.8', '3.9', '3.10']
runs-on: ${{ matrix.os }}
steps:
- uses: actions/checkout@v3
- name: Setup Python ${{ matrix.python-version }}
uses: actions/setup-python@v4
with:
python-version: ${{ matrix.python-version }}
- name: Build and test
run: |
pip install pytest
mkdir build && cd build
cmake ..
make
pytest ../tests/
import pytest
import my_module
class TestRegression:
"""回归测试:确保历史 bug 不再出现"""
def test_issue_123_memory_leak(self):
"""修复 #123:内存泄漏问题"""
import my_module
for _ in range(1000):
obj = my_module.create_object()
del obj
def test_issue_456_crash(self):
"""修复 #456:特定输入导致崩溃"""
result = my_module.process([1, 2, 3])
assert result is not None
def test_issue_789_wrong_result(self):
"""修复 #789:计算结果错误"""
result = my_module.compute(0.1, 0.2)
assert abs(result - 0.3) < 1e-10

关键洞察:CI 是质量保障的关键。每次 PR 都应运行完整的测试套件。使用矩阵测试覆盖多个 Python 版本和操作系统。

性能测试确保优化不会破坏功能,同时验证性能目标。

import timeit
import my_module
class TestPerformance:
def test_vector_sum_performance(self):
"""向量求和性能测试"""
data = list(range(100000))
def run():
return my_module.sum(data)
# 运行基准测试
number = 100
time_taken = timeit.timeit(run, number=number)
avg_time = time_taken / number
print(f"Average time: {avg_time*1000:.3f}ms")
# 性能断言
assert avg_time < 0.1, f"Performance regression: {avg_time*1000:.3f}ms > 100ms"
def test_matrix_multiply_performance(self):
"""矩阵乘法性能测试"""
size = 500
a = [[1.0] * size for _ in range(size)]
b = [[2.0] * size for _ in range(size)]
time_taken = timeit.timeit(
lambda: my_module.matrix_multiply(a, b),
number=10
)
print(f"Matrix multiply time: {time_taken:.3f}s")
assert time_taken < 5.0, "Performance regression detected"
import subprocess
import sys
def test_cpu_performance():
"""使用 perf 测量 CPU 性能"""
code = """
import my_module
data = list(range(1000000))
result = my_module.sum(data)
"""
result = subprocess.run(
["perf", "stat", "-e", "cycles,instructions", "python3", "-c", code],
capture_output=True,
text=True
)
print(result.stdout)
# 检查性能指标
assert "cycles" in result.stdout
import time
import my_module
BENCHMARK_HISTORY = {
"sum_1m": 0.05, # 50ms
"matrix_500": 1.0, # 1s
"sort_100k": 0.2, # 200ms
}
def check_performance_regression():
"""检查性能是否退化"""
print("Running performance regression checks...")
# 向量求和
data = list(range(1000000))
start = time.perf_counter()
my_module.sum(data)
elapsed = time.perf_counter() - start
threshold = BENCHMARK_HISTORY["sum_1m"] * 1.2 # 20% 容限
if elapsed > threshold:
print(f"REGRESSION: sum_1m took {elapsed:.3f}s > {threshold:.3f}s")
return False
print("All performance checks passed")
return True

关键洞察:性能基准测试应该在 CI 中持续运行,及时发现性能退化。设置合理的阈值,允许一定波动但捕捉显著退化。

测试策略总结:

测试类型工具覆盖范围
Python 接口pytestPython 绑定层
C++ 逻辑gtest核心算法
性能pytest-benchmark性能回归
覆盖率coverage.py + gcov测试盲区
CIGitHub Actions多平台多版本

实战建议:测试金字塔,底层是大量快速的单元测试,上层是较少的集成测试。Python 端测试绑定正确性,C++ 端测试核心逻辑,两者缺一不可。