You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Rust中测试经PyO3封装的Python暴露函数?

在Rust中测试PyO3封装的Python函数

1. 基于GIL的Rust测试框架

PyO3的测试必须在Python全局解释器锁(GIL)上下文里执行,你可以直接用Rust的#[test]宏结合Python::with_gil实现:

use pyo3::prelude::*;

// 示例:PyO3暴露的函数
#[pyfunction]
fn add_one(x: i32) -> PyResult<i32> {
    Ok(x + 1)
}

#[test]
fn test_add_one() {
    Python::with_gil(|py| {
        // 调用封装函数,传入常规参数
        let result = add_one(py, 5).unwrap();
        assert_eq!(result, 6);

        // 测试边界值
        let result_neg = add_one(py, -10).unwrap();
        assert_eq!(result_neg, -9);
    });
}

2. 验证Python原生类型处理

针对接收Python字符串、列表、字典等原生类型的封装函数,直接构造对应PyO3类型实例传入测试:

#[pyfunction]
fn count_chars(s: &PyString) -> PyResult<usize> {
    Ok(s.to_str()?.chars().count())
}

#[test]
fn test_count_chars() {
    Python::with_gil(|py| {
        // 构造带非ASCII字符的Python字符串
        let py_str = PyString::new(py, "Hello 世界");
        let count = count_chars(py, py_str).unwrap();
        assert_eq!(count, 7);

        // 测试空字符串场景
        let empty_str = PyString::new(py, "");
        let empty_count = count_chars(py, empty_str).unwrap();
        assert_eq!(empty_count, 0);
    });
}

处理列表类型时,可直接从Rust数据结构转换:

#[pyfunction]
fn sum_list(lst: &PyList) -> PyResult<i32> {
    let mut total = 0;
    for item in lst.iter() {
        total += item.extract::<i32>()?;
    }
    Ok(total)
}

#[test]
fn test_sum_list() {
    Python::with_gil(|py| {
        let py_list = PyList::new(py, &[1, 2, 3, 4]);
        let sum = sum_list(py, py_list).unwrap();
        assert_eq!(sum, 10);

        // 测试类型不匹配的错误场景
        let invalid_list = PyList::new(py, &[1, "two", 3]);
        assert!(sum_list(py, invalid_list).is_err());
        assert!(PyErr::occurred(py).is_some());
    });
}

3. 测试异常处理逻辑

主动触发错误场景,验证封装函数抛出的Python异常是否符合预期:

#[pyfunction]
fn divide(a: f64, b: f64) -> PyResult<f64> {
    if b == 0.0 {
        return Err(PyErr::new::<pyo3::exceptions::ZeroDivisionError, _>(
            "Cannot divide by zero",
        ));
    }
    Ok(a / b)
}

#[test]
fn test_divide_errors() {
    Python::with_gil(|py| {
        let err = divide(py, 10.0, 0.0).unwrap_err();
        // 验证异常类型
        assert!(err.is_instance_of::<pyo3::exceptions::ZeroDivisionError>(py));
        
        // 验证错误消息
        let msg = err.value(py).to_string();
        assert!(msg.contains("Cannot divide by zero"));
    });
}

4. 简化测试的工具(可选)

如果觉得手动处理GIL样板代码麻烦,可以使用pyo3-test crate,它的#[pyo3_test]宏会自动处理GIL初始化:

先在Cargo.toml添加依赖:

[dev-dependencies]
pyo3-test = "0.21"

再编写简化的测试:

use pyo3_test::pyo3_test;

#[pyo3_test]
fn test_simplified_add_one(py: Python<'_>) {
    let result = add_one(py, 5).unwrap();
    assert_eq!(result, 6);
}

内容的提问来源于stack exchange,提问作者MusicalNinja

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.27 05:12:38