Skip to content

Commit 05f3668

Browse files
committed
use write_attrs to patch up the CUDA let_(value|error|stopped) tests
1 parent 921cd81 commit 05f3668

4 files changed

Lines changed: 56 additions & 26 deletions

File tree

include/exec/env.hpp

Lines changed: 17 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -269,23 +269,36 @@ namespace experimental::execution
269269
{
270270
template <class _Sender, class _Attrs>
271271
STDEXEC_ATTRIBUTE(host, device)
272-
constexpr auto operator()(_Sender snd, _Attrs __attrs_) const -> __sender<_Sender, _Attrs>
272+
constexpr auto operator()(_Sender snd, _Attrs __attrs) const -> __sender<_Sender, _Attrs>
273273
{
274274
return __sender<_Sender, _Attrs>{static_cast<_Sender&&>(snd),
275-
static_cast<_Attrs&&>(__attrs_)};
275+
static_cast<_Attrs&&>(__attrs)};
276276
}
277277

278278
template <class _Attrs>
279279
STDEXEC_ATTRIBUTE(host, device)
280-
constexpr auto operator()(_Attrs __attrs_) const
280+
constexpr auto operator()(_Attrs __attrs) const
281281
{
282-
return STDEXEC::__closure(*this, static_cast<_Attrs&&>(__attrs_));
282+
return STDEXEC::__closure(*this, static_cast<_Attrs&&>(__attrs));
283283
}
284284
};
285285
} // namespace __write_attrs
286286

287287
inline constexpr __write_attrs::__write_attrs_t write_attrs{};
288288

289+
template <class Tag = STDEXEC::set_value_t, STDEXEC::scheduler Scheduler>
290+
STDEXEC_ATTRIBUTE(host, device)
291+
constexpr auto completes_on(Scheduler sched) noexcept
292+
{
293+
return exec::write_attrs(STDEXEC::prop{STDEXEC::get_completion_scheduler<Tag>, sched});
294+
}
295+
296+
template <class Tag = STDEXEC::set_value_t, STDEXEC::sender Sender, STDEXEC::scheduler Scheduler>
297+
STDEXEC_ATTRIBUTE(host, device)
298+
constexpr auto completes_on(Sender sndr, Scheduler sched) noexcept
299+
{
300+
return exec::completes_on<Tag>(sched)(std::move(sndr));
301+
}
289302
} // namespace experimental::execution
290303

291304
namespace exec = experimental::execution;

test/nvexec/let_error.cpp

Lines changed: 15 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,8 @@ namespace
1616
{
1717
nvexec::stream_context stream_ctx{};
1818

19-
auto snd = ex::just_error(42) | ex::continues_on(stream_ctx.get_scheduler())
19+
auto snd = ex::just_error(42) //
20+
| ex::continues_on(stream_ctx.get_scheduler()) //
2021
| ex::let_error([](int) { return ex::just(); });
2122
STATIC_REQUIRE(ex::sender<decltype(snd)>);
2223
(void) snd;
@@ -29,7 +30,8 @@ namespace
2930
flags_storage_t flags_storage{};
3031
auto flags = flags_storage.get();
3132

32-
auto snd = ex::just_error(42) | ex::continues_on(stream_ctx.get_scheduler())
33+
auto snd = ex::just_error(42) //
34+
| ex::continues_on(stream_ctx.get_scheduler()) //
3335
| ex::let_error(
3436
[=](int err)
3537
{
@@ -38,9 +40,9 @@ namespace
3840
flags.set();
3941
}
4042

41-
return ex::just()
42-
| exec::write_attrs(ex::prop{ex::get_domain, nvexec::stream_domain()});
43-
});
43+
return ex::just();
44+
})
45+
| exec::completes_on(stream_ctx.get_scheduler());
4446
STDEXEC::sync_wait(std::move(snd));
4547

4648
REQUIRE(flags_storage.all_set_once());
@@ -54,7 +56,8 @@ namespace
5456
flags_storage_t<2> flags_storage{};
5557
auto flags = flags_storage.get();
5658

57-
auto snd = ex::just_error(42) | ex::continues_on(stream_ctx.get_scheduler())
59+
auto snd = ex::just_error(42) //
60+
| ex::continues_on(stream_ctx.get_scheduler()) //
5861
| ex::let_error(
5962
[flags](int err)
6063
{
@@ -63,9 +66,9 @@ namespace
6366
flags.set(0);
6467
}
6568

66-
return ex::just()
67-
| exec::write_attrs(ex::prop{ex::get_domain, nvexec::stream_domain()});
69+
return ex::just();
6870
})
71+
| exec::completes_on(stream_ctx.get_scheduler())
6972
| a_sender(
7073
[flags]
7174
{
@@ -86,7 +89,8 @@ namespace
8689
flags_storage_t flags_storage{};
8790
auto flags = flags_storage.get();
8891

89-
auto snd = ex::just_error(42) | ex::continues_on(stream_ctx.get_scheduler())
92+
auto snd = ex::just_error(42) //
93+
| ex::continues_on(stream_ctx.get_scheduler()) //
9094
| a_sender([]() noexcept {})
9195
| ex::let_error(
9296
[=](int err)
@@ -97,7 +101,8 @@ namespace
97101
}
98102

99103
return ex::schedule(sch);
100-
});
104+
})
105+
| exec::completes_on(stream_ctx.get_scheduler());
101106
STDEXEC::sync_wait(std::move(snd));
102107

103108
REQUIRE(flags_storage.all_set_once());

test/nvexec/let_stopped.cpp

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -38,7 +38,8 @@ namespace
3838
}
3939

4040
return ex::just();
41-
});
41+
})
42+
| exec::completes_on(stream_ctx.get_scheduler());
4243
STDEXEC::sync_wait(std::move(snd));
4344

4445
REQUIRE(flags_storage.all_set_once());
@@ -52,7 +53,8 @@ namespace
5253
flags_storage_t<2> flags_storage{};
5354
auto flags = flags_storage.get();
5455

55-
auto snd = ex::just_stopped() | ex::continues_on(stream_ctx.get_scheduler())
56+
auto snd = ex::just_stopped() //
57+
| ex::continues_on(stream_ctx.get_scheduler()) //
5658
| ex::let_stopped(
5759
[flags]
5860
{
@@ -63,6 +65,7 @@ namespace
6365

6466
return ex::just();
6567
})
68+
| exec::completes_on(stream_ctx.get_scheduler())
6669
| a_sender(
6770
[flags]
6871
{
@@ -83,7 +86,9 @@ namespace
8386
flags_storage_t flags_storage{};
8487
auto flags = flags_storage.get();
8588

86-
auto snd = ex::just_stopped() | ex::continues_on(sch) | a_sender([]() noexcept {})
89+
auto snd = ex::just_stopped() //
90+
| ex::continues_on(sch) //
91+
| a_sender([]() noexcept {}) //
8792
| ex::let_stopped(
8893
[=]
8994
{
@@ -93,7 +98,8 @@ namespace
9398
}
9499

95100
return ex::schedule(sch);
96-
});
101+
})
102+
| exec::completes_on(sch);
97103

98104
STDEXEC::sync_wait(std::move(snd));
99105

test/nvexec/let_value.cpp

Lines changed: 14 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,8 @@ namespace
3939
flags.set();
4040
}
4141
return ex::just();
42-
});
42+
})
43+
| exec::completes_on(stream_ctx.get_scheduler());
4344
ex::sync_wait(std::move(snd));
4445

4546
REQUIRE(flags_storage.all_set_once());
@@ -64,7 +65,8 @@ namespace
6465
}
6566
}
6667
return ex::just();
67-
});
68+
})
69+
| exec::completes_on(stream_ctx.get_scheduler());
6870
ex::sync_wait(std::move(snd));
6971

7072
REQUIRE(flags_storage.all_set_once());
@@ -90,7 +92,8 @@ namespace
9092
}
9193
}
9294
return ex::just();
93-
});
95+
})
96+
| exec::completes_on(stream_ctx.get_scheduler());
9497
ex::sync_wait(std::move(snd));
9598

9699
REQUIRE(flags_storage.all_set_once());
@@ -101,7 +104,8 @@ namespace
101104
nvexec::stream_context stream_ctx{};
102105

103106
auto snd = ex::schedule(stream_ctx.get_scheduler())
104-
| ex::let_value([=]() { return ex::just(is_on_gpu()); });
107+
| ex::let_value([=]() { return ex::just(is_on_gpu()); })
108+
| exec::completes_on(stream_ctx.get_scheduler());
105109
auto const [result] = ex::sync_wait(std::move(snd)).value();
106110

107111
REQUIRE(result == 1);
@@ -126,6 +130,7 @@ namespace
126130

127131
return ex::just();
128132
})
133+
| exec::completes_on(stream_ctx.get_scheduler())
129134
| a_sender(
130135
[flags]
131136
{
@@ -156,7 +161,8 @@ namespace
156161
}
157162

158163
return ex::schedule(sch);
159-
});
164+
})
165+
| exec::completes_on(sch);
160166
ex::sync_wait(std::move(snd));
161167

162168
REQUIRE(flags_storage.all_set_once());
@@ -169,9 +175,9 @@ namespace
169175
flags_storage_t flags_storage{};
170176
auto flags = flags_storage.get();
171177

172-
auto snd = ex::schedule(sch) //
173-
| ex::let_value([] { return nvexec::get_stream(); }) //
174-
| exec::write_attrs(ex::prop{ex::get_completion_scheduler<ex::set_value_t>, sch}) //
178+
auto snd = ex::schedule(sch) //
179+
| ex::let_value([] { return nvexec::get_stream(); }) //
180+
| exec::completes_on(sch) //
175181
| ex::then(
176182
[flags](cudaStream_t stream)
177183
{

0 commit comments

Comments
 (0)