diff --git a/src/pybind11-utils/functional.cpp b/src/pybind11-utils/functional.cpp new file mode 100644 index 0000000..4128885 --- /dev/null +++ b/src/pybind11-utils/functional.cpp @@ -0,0 +1,27 @@ +#include "pybind11_utils/functional.h" + +namespace py = pybind11; + +namespace mo2::python::detail { + + bool has_compatible_arity(py::function fn, std::size_t arity) + { + auto inspect = py::module_::import("inspect"); + auto arg_spec = inspect.attr("getfullargspec")(fn); + py::object args = arg_spec.attr("args"), varargs = arg_spec.attr("varargs"), + defaults = arg_spec.attr("defaults"); + + auto args_count = args.is(py::none()) ? 0 : py::len(args); + auto defaults_count = defaults.is(py::none()) ? 0 : py::len(defaults); + + if (inspect.attr("ismethod")(fn).cast() && py::hasattr(fn, "__self__")) { + --args_count; + } + + auto required_count = args_count - defaults_count; + + return required_count <= arity // cannot require more parameters than given, + && (args_count >= arity || varargs); // must accept enough parameters. + } + +} // namespace mo2::python::detail diff --git a/src/pybind11-utils/include/pybind11_utils/functional.h b/src/pybind11-utils/include/pybind11_utils/functional.h index d23be5a..67c599e 100644 --- a/src/pybind11-utils/include/pybind11_utils/functional.h +++ b/src/pybind11-utils/include/pybind11_utils/functional.h @@ -1,7 +1,158 @@ #ifndef PYTHON_PYBIND11_FUNCTIONAL_H #define PYTHON_PYBIND11_FUNCTIONAL_H -// TODO -#include +#include + +namespace mo2::python::detail { + + // check if the given function is valid for a C++ function with the + // given arity + // + bool has_compatible_arity(pybind11::function handle, std::size_t arity); + +} // namespace mo2::python::detail + +namespace pybind11::detail { + + // custom type_caster for std::function<> + // + // most of this is from pybind11 except that we also check arity of the function to + // allow overloaded function based on arity of argument + // + template + struct type_caster> { + using type = std::function; + using retval_type = + conditional_t::value, void_type, Return>; + using function_type = Return (*)(Args...); + + public: + bool load(handle src, bool convert) + { + if (src.is_none()) { + // Defer accepting None to other overloads (if we aren't in convert + // mode): + if (!convert) { + return false; + } + return true; + } + + if (!isinstance(src)) { + return false; + } + + auto func = reinterpret_borrow(src); + + /* + When passing a C++ function as an argument to another C++ + function via Python, every function call would normally involve + a full C++ -> Python -> C++ roundtrip, which can be prohibitive. + Here, we try to at least detect the case where the function is + stateless (i.e. function pointer or lambda function without + captured variables), in which case the roundtrip can be avoided. + */ + if (auto cfunc = func.cpp_function()) { + auto* cfunc_self = PyCFunction_GET_SELF(cfunc.ptr()); + if (isinstance(cfunc_self)) { + auto c = reinterpret_borrow(cfunc_self); + auto* rec = (function_record*)c; + + while (rec != nullptr) { + if (rec->is_stateless && + same_type(typeid(function_type), + *reinterpret_cast( + rec->data[1]))) { + struct capture { + function_type f; + }; + value = ((capture*)&rec->data)->f; + return true; + } + rec = rec->next; + } + } + // PYPY segfaults here when passing builtin function like sum. + // Raising an fail exception here works to prevent the segfault, but + // only on gcc. See PR #1413 for full details + } + + // !MO2! - check arity + + if (!mo2::python::detail::has_compatible_arity(func, sizeof...(Args))) { + return false; + } + + // !MO2! - everything below is copy/paste from pybind11 + + // ensure GIL is held during functor destruction + struct func_handle { + function f; +#if !(defined(_MSC_VER) && _MSC_VER == 1916 && defined(PYBIND11_CPP17)) + // This triggers a syntax error under very special conditions (very + // weird indeed). + explicit +#endif + func_handle(function&& f_) noexcept + : f(std::move(f_)) + { + } + func_handle(const func_handle& f_) + { + operator=(f_); + } + func_handle& operator=(const func_handle& f_) + { + gil_scoped_acquire acq; + f = f_.f; + return *this; + } + ~func_handle() + { + gil_scoped_acquire acq; + function kill_f(std::move(f)); + } + }; + + // to emulate 'move initialization capture' in C++11 + struct func_wrapper { + func_handle hfunc; + explicit func_wrapper(func_handle&& hf) noexcept : hfunc(std::move(hf)) + { + } + Return operator()(Args... args) const + { + gil_scoped_acquire acq; + object retval(hfunc.f(std::forward(args)...)); + return retval.template cast(); + } + }; + + value = func_wrapper(func_handle(std::move(func))); + return true; + } + + template + static handle cast(Func&& f_, return_value_policy policy, handle /* parent */) + { + if (!f_) { + return none().inc_ref(); + } + + auto result = f_.template target(); + if (result) { + return cpp_function(*result, policy).release(); + } + return cpp_function(std::forward(f_), policy).release(); + } + + PYBIND11_TYPE_CASTER(type, const_name("Callable[[") + + concat(make_caster::name...) + + const_name("], ") + + make_caster::name + + const_name("]")); + }; + +} // namespace pybind11::detail #endif diff --git a/tests/python/test_functional.cpp b/tests/python/test_functional.cpp new file mode 100644 index 0000000..4858516 --- /dev/null +++ b/tests/python/test_functional.cpp @@ -0,0 +1,38 @@ +#include "pybind11_utils/functional.h" + +#include + +PYBIND11_MODULE(functional, m) +{ + m.def("fn_0_arg", [](std::function const& fn) { + return fn(); + }); + + m.def("fn_1_arg", [](std::function const& fn, int a) { + return fn(a); + }); + + m.def("fn_2_arg", [](std::function const& fn, int a, int b) { + return fn(a, b); + }); + + m.def("fn_0_or_1_arg", [](std::function const& fn) { + return fn(); + }); + + m.def("fn_0_or_1_arg", [](std::function const& fn) { + return fn(1); + }); + + m.def("fn_1_or_2_or_3_arg", [](std::function const& fn) { + return fn(1); + }); + + m.def("fn_1_or_2_or_3_arg", [](std::function const& fn) { + return fn(1, 2); + }); + + m.def("fn_1_or_2_or_3_arg", [](std::function const& fn) { + return fn(1, 2, 3); + }); +} diff --git a/tests/python/test_functional.py b/tests/python/test_functional.py new file mode 100644 index 0000000..a5cb30a --- /dev/null +++ b/tests/python/test_functional.py @@ -0,0 +1,43 @@ +import mobase +import pytest + +m = pytest.importorskip("mobase_tests.functional") + + +def test_guessed_string(): + + # available functions: + # - fn_0_arg, fn_1_arg, fn_2_arg + # - fn_0_or_1_arg, fn_1_or_2_or_3_arg + + assert m.fn_0_arg(lambda: 0) == 0 + assert m.fn_0_arg(lambda: 5) == 5 + assert m.fn_0_arg(lambda x=2: x) == 2 + assert m.fn_0_arg(lambda *args: len(args)) == 0 + + assert m.fn_1_arg(lambda x: x, 4) == 4 + assert m.fn_1_arg(lambda *args: sum(args), 8) == 8 + assert m.fn_1_arg(lambda x=2, y=4: x + y, 3) == 7 + assert m.fn_1_arg(lambda x, *args, **kwargs: x + len(args), 5) == 5 + + assert m.fn_2_arg(lambda x, y: x * y, 4, 5) == 20 + assert m.fn_2_arg(lambda x, y=3: x * y, 4, 2) == 8 + assert m.fn_2_arg(lambda *args: sum(args), 8, 9) == 17 + assert m.fn_2_arg(lambda x=2, y=4: x + y, 3, 3) == 6 + assert m.fn_2_arg(lambda x, *args, **kwargs: x + len(args), 5, 2) == 6 + + assert m.fn_0_or_1_arg(lambda: 3) == 3 + assert m.fn_0_or_1_arg(lambda x: x) == 1 + + # the 0 arg is bound first, both are possible, the 0 arg is chosen + assert m.fn_0_or_1_arg(lambda x=2: x) == 2 + + with pytest.raises(TypeError): + m.fn_1_or_2_or_3_arg(lambda: 0) + + assert m.fn_1_or_2_or_3_arg(lambda x=4: x) == 1 + assert m.fn_1_or_2_or_3_arg(lambda x, y: x + y) == 3 # 1 + 2 + assert m.fn_1_or_2_or_3_arg(lambda x, y, z: x * y * z) == 6 # 1 * 2 * 3 + + # the 1 arg is bound first + assert m.fn_1_or_2_or_3_arg(lambda *args: sum(args)) == 1