diff --git a/include/nvexec/stream/launch.cuh b/include/nvexec/stream/launch.cuh index 187da083b..37b1a1881 100644 --- a/include/nvexec/stream/launch.cuh +++ b/include/nvexec/stream/launch.cuh @@ -141,7 +141,11 @@ namespace nv::execution static_cast(self).sndr_, static_cast(rcvr), [&](_strm::opstate_base& stream_provider) -> receiver_t - { return receiver_t(stream_provider, self.fun_, self.params_); }); + { + return receiver_t(stream_provider, + static_cast(self).fun_, + self.params_); + }); } STDEXEC_EXPLICIT_THIS_END(connect) diff --git a/test/nvexec/launch.cpp b/test/nvexec/launch.cpp index 7a9d69a98..54cbec6c5 100644 --- a/test/nvexec/launch.cpp +++ b/test/nvexec/launch.cpp @@ -18,6 +18,8 @@ #include #include +#include + #include "common.cuh" #include "nvexec/stream_context.cuh" @@ -50,6 +52,31 @@ namespace return std::accumulate(input.begin(), input.end(), 0); } + struct move_only_launch_handler + { + move_only_launch_handler() = default; + move_only_launch_handler(move_only_launch_handler const &) = delete; + + STDEXEC_ATTRIBUTE(host, device) + move_only_launch_handler(move_only_launch_handler&&) = default; + + STDEXEC_ATTRIBUTE(host, device) void operator()(cudaStream_t) const {} + }; + + static_assert(std::is_trivially_copyable_v); + static_assert(!std::is_copy_constructible_v); + + TEST_CASE("nvexec launch supports move-only function objects", + "[cuda][stream][adaptors][launch]") + { + nvexec::stream_context stream_ctx{}; + + auto snd = STDEXEC::just() | STDEXEC::continues_on(stream_ctx.get_scheduler()) + | nvexec::launch(move_only_launch_handler{}); + + REQUIRE(STDEXEC::sync_wait(std::move(snd)).has_value()); + } + TEST_CASE("nvexec launch advertises CUDA launch errors", "[cuda][stream][adaptors][launch]") { nvexec::stream_context stream{};